Merge main and host authorization-code OAuth in core

This commit is contained in:
leokun
2026-09-04 17:16:29 +08:00
224 changed files with 10351 additions and 21466 deletions
@@ -0,0 +1,20 @@
/** Official Google Antigravity OAuth client configuration. */
const _P1 = "1071006060591";
const _P2 = "tmhssin2h21lcre235vtolojh4g403ep";
const _P3 = "apps.googleusercontent.com";
export const CLIENT_ID = [_P1, _P2, _P3].join("-").replace("-apps", ".apps");
const _S1 = "GOCSPX";
const _S2 = "K58FWR486LdLJ1mLB8sXC4z6qDAf";
export const CLIENT_SECRET = [_S1, _S2].join("-");
export const SCOPES = [
"https://www.googleapis.com/auth/cloud-platform",
"https://www.googleapis.com/auth/userinfo.email",
"https://www.googleapis.com/auth/userinfo.profile",
"https://www.googleapis.com/auth/cclog",
"https://www.googleapis.com/auth/experimentsandconfigs",
];
export const GOOGLE_AUTHORIZATION_URL = "https://accounts.google.com/o/oauth2/v2/auth";
export const GOOGLE_TOKEN_URL = "https://oauth2.googleapis.com/token";
@@ -1,12 +1,7 @@
import { defineProviderPlugin } from "cursor-byok:plugin";
import { antigravityDeviceOAuth } from "./oauth.ts";
import { antigravityAuthorizationCodeOAuth } from "./oauth.ts";
import { antigravityProvider } from "./provider.ts";
import {
credentialImport,
presentAccount,
refreshAccount,
RESOURCE_TYPE,
} from "./resources.ts";
import { credentialImport, presentAccount, refreshAccount, RESOURCE_TYPE } from "./resources.ts";
export default defineProviderPlugin({
providers: [antigravityProvider],
@@ -16,7 +11,7 @@ export default defineProviderPlugin({
"en-US": "Google accounts & API keys",
"zh-CN": "Google 账号与 API 密钥",
},
add: [antigravityDeviceOAuth],
add: [antigravityAuthorizationCodeOAuth],
import: credentialImport,
present: presentAccount,
refresh: refreshAccount,
@@ -277,7 +277,9 @@ export function parseAntigravityModels(payload: unknown): ModelDefinition[] {
models.push({
id,
displayName,
capabilities: { images: model.supportsImages === true || id.includes("gemini") || id.includes("claude") },
capabilities: {
images: model.supportsImages === true || id.includes("gemini") || id.includes("claude"),
},
maxOutputTokens,
privateData: { reasoningEfforts },
});
@@ -318,16 +320,19 @@ export const antigravityModels: ModelSupport = {
for (const endpoint of ANTIGRAVITY_ENDPOINTS) {
for (const bodyPayload of payloads) {
try {
const response = await context.network.fetch(`${endpoint}${FETCH_AVAILABLE_MODELS_PATH}`, {
method: "POST",
headers: {
authorization: `Bearer ${data.accessToken}`,
"content-type": "application/json",
"user-agent": ANTIGRAVITY_USER_AGENT,
...ANTIGRAVITY_CLIENT_HEADERS,
const response = await context.network.fetch(
`${endpoint}${FETCH_AVAILABLE_MODELS_PATH}`,
{
method: "POST",
headers: {
authorization: `Bearer ${data.accessToken}`,
"content-type": "application/json",
"user-agent": ANTIGRAVITY_USER_AGENT,
...ANTIGRAVITY_CLIENT_HEADERS,
},
body: bodyPayload,
},
body: bodyPayload,
});
);
if (response.status >= 200 && response.status < 300) {
const body = JSON.parse(response.body);
const models = parseAntigravityModels(body);
+97 -139
View File
@@ -1,38 +1,24 @@
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
import type { OAuth2AddMethod, OAuth2Begin, OAuth2Poll } from "cursor-byok:resource";
import type {
OAuth2AuthorizationCodeAddMethod,
OAuth2AuthorizationCodeBegin,
ResourceDraft,
} from "cursor-byok:resource";
import { credentialDraft, queryAccountQuota } from "./resources.ts";
/**
* Official Google Antigravity OAuth Client credentials.
*/
const _P1 = "1071006060591";
const _P2 = "tmhssin2h21lcre235vtolojh4g403ep";
const _P3 = "apps.googleusercontent.com";
export const CLIENT_ID = [_P1, _P2, _P3].join("-").replace("-apps", ".apps");
const _S1 = "GOCSPX";
const _S2 = "K58FWR486LdLJ1mLB8sXC4z6qDAf";
export const CLIENT_SECRET = [_S1, _S2].join("-");
import {
CLIENT_ID,
CLIENT_SECRET,
GOOGLE_AUTHORIZATION_URL,
GOOGLE_TOKEN_URL,
SCOPES,
} from "./google_oauth.ts";
export const CALLBACK_PORT = 51121;
export const CALLBACK_PATH = "/oauth-callback";
export const REDIRECT_URI = `http://127.0.0.1:${CALLBACK_PORT}${CALLBACK_PATH}`;
export const SCOPES = [
"https://www.googleapis.com/auth/cloud-platform",
"https://www.googleapis.com/auth/userinfo.email",
"https://www.googleapis.com/auth/userinfo.profile",
"https://www.googleapis.com/auth/cclog",
"https://www.googleapis.com/auth/experimentsandconfigs",
];
const AUTHORIZATION_LIFETIME_MS = 5 * 60 * 1000;
const AUTH_URL = "https://accounts.google.com/o/oauth2/v2/auth";
const TOKEN_URL = "https://oauth2.googleapis.com/token";
type Session = {
state: string;
createdAt: number;
};
type Session = { createdAtMs: number };
function object(value: unknown): Record<string, unknown> | null {
return value !== null && typeof value === "object" && !Array.isArray(value)
@@ -54,146 +40,118 @@ function parseBody(body: string): Record<string, unknown> {
function parseSession(value: JsonValue): Session {
const session = object(value);
const state = text(session?.state);
const createdAt = typeof session?.createdAt === "number" ? session.createdAt : Date.now();
if (!state) throw new Error("Antigravity OAuth session is invalid");
return { state, createdAt };
const createdAtMs = session?.createdAtMs;
if (typeof createdAtMs !== "number") throw new Error("Antigravity OAuth session is invalid");
if (Date.now() - createdAtMs > AUTHORIZATION_LIFETIME_MS) {
throw new Error("Google authorization expired. Please try again.");
}
return { createdAtMs };
}
function randomState(): string {
const array = new Uint8Array(24);
crypto.getRandomValues(array);
return Array.from(array, (byte) => byte.toString(16).padStart(2, "0")).join("");
}
async function begin(_context: PluginContext): Promise<OAuth2Begin> {
const state = randomState();
async function begin(
input: { redirectUri: string; state: string; codeChallenge: string },
_context: PluginContext,
): Promise<OAuth2AuthorizationCodeBegin> {
const authParams = new URLSearchParams({
client_id: CLIENT_ID,
response_type: "code",
redirect_uri: REDIRECT_URI,
redirect_uri: input.redirectUri,
scope: SCOPES.join(" "),
state,
state: input.state,
code_challenge: input.codeChallenge,
code_challenge_method: "S256",
access_type: "offline",
prompt: "consent",
});
const verificationUrl = `${AUTH_URL}?${authParams.toString()}`;
const session: Session = { state, createdAt: Date.now() };
return {
session: session as unknown as JsonValue,
userCode: "Google Sign-in",
verificationUrl,
verificationUrlComplete: verificationUrl,
expiresAtMs: Date.now() + 300 * 1000,
pollIntervalMs: 1500,
session: { createdAtMs: Date.now() },
authorizationUrl: `${GOOGLE_AUTHORIZATION_URL}?${authParams.toString()}`,
expiresAtMs: Date.now() + AUTHORIZATION_LIFETIME_MS,
};
}
async function poll(sessionValue: JsonValue, context: PluginContext): Promise<OAuth2Poll> {
const session = parseSession(sessionValue);
// Check if authorization timed out
if (Date.now() - session.createdAt > 300 * 1000) {
return { status: "failed", message: "Sign-in timed out. Please try again." };
async function complete(
sessionValue: JsonValue,
input: { code: string; redirectUri: string; codeVerifier: string },
context: PluginContext,
): Promise<ResourceDraft[]> {
parseSession(sessionValue);
const tokenResponse = await context.network.fetch(GOOGLE_TOKEN_URL, {
method: "POST",
headers: {
accept: "application/json",
"content-type": "application/x-www-form-urlencoded",
},
body: new URLSearchParams({
client_id: CLIENT_ID,
client_secret: CLIENT_SECRET,
code: input.code,
code_verifier: input.codeVerifier,
grant_type: "authorization_code",
redirect_uri: input.redirectUri,
}).toString(),
});
const tokenBody = parseBody(tokenResponse.body);
if (tokenResponse.status < 200 || tokenResponse.status >= 300) {
const detail = text(tokenBody.error_description) ?? text(tokenBody.error) ??
`HTTP ${tokenResponse.status}`;
throw new Error(`Google token exchange failed: ${detail}`);
}
// Attempt to check if local callback server on 51121 received the auth code
const accessToken = text(tokenBody.access_token);
if (!accessToken) throw new Error("Google token response did not include an access token");
let displayName = text(tokenBody.email);
try {
const callbackCheck = await context.network.fetch(
`http://127.0.0.1:${CALLBACK_PORT}/auth-status?state=${session.state}`,
{ method: "GET" },
const userInfoResponse = await context.network.fetch(
"https://www.googleapis.com/oauth2/v1/userinfo?alt=json",
{
method: "GET",
headers: { authorization: `Bearer ${accessToken}`, accept: "application/json" },
},
);
if (callbackCheck.status === 200) {
const body = parseBody(callbackCheck.body);
const code = text(body.code);
if (code) {
// Exchange code for tokens
const tokenResponse = await context.network.fetch(TOKEN_URL, {
method: "POST",
headers: {
accept: "application/json",
"content-type": "application/x-www-form-urlencoded",
},
body: new URLSearchParams({
client_id: CLIENT_ID,
client_secret: CLIENT_SECRET,
code,
grant_type: "authorization_code",
redirect_uri: REDIRECT_URI,
}).toString(),
});
const tokenBody = parseBody(tokenResponse.body);
if (tokenResponse.status >= 200 && tokenResponse.status < 300) {
const accessToken = text(tokenBody.access_token);
if (accessToken) {
let email: string | null = text(tokenBody.email);
try {
const userInfoRes = await context.network.fetch(
"https://www.googleapis.com/oauth2/v1/userinfo?alt=json",
{
method: "GET",
headers: {
authorization: `Bearer ${accessToken}`,
accept: "application/json",
},
},
);
if (userInfoRes.status === 200) {
const userInfo = parseBody(userInfoRes.body);
email = text(userInfo.email) ?? email;
}
} catch {
// Ignore error, fallback to default display name
}
let projectId = "bamboo-precept-lgxtn";
let quota = null;
try {
const res = await queryAccountQuota(accessToken, context.network);
projectId = res.projectId;
quota = res.quota;
} catch {
// Ignore error
}
return {
status: "completed",
resources: [
await credentialDraft({
accessToken,
refreshToken: text(tokenBody.refresh_token),
displayName: email ?? "Google Antigravity",
projectId,
quota,
}),
],
};
}
} else {
const errMsg = text(tokenBody.error_description ?? tokenBody.error) ?? `HTTP ${tokenResponse.status}`;
return { status: "failed", message: `Token exchange failed: ${errMsg}` };
}
}
if (userInfoResponse.status >= 200 && userInfoResponse.status < 300) {
displayName = text(parseBody(userInfoResponse.body).email) ?? displayName;
}
} catch {
// Network retry on pending callback
// Account identity has a token fingerprint fallback.
}
return { status: "pending" };
let projectId = "bamboo-precept-lgxtn";
let quota = null;
try {
const result = await queryAccountQuota(accessToken, context.network);
projectId = result.projectId;
quota = result.quota;
} catch {
// Quota can be refreshed after the account has been persisted.
}
const expiresIn = typeof tokenBody.expires_in === "number" ? tokenBody.expires_in : null;
return [
await credentialDraft({
accessToken,
refreshToken: text(tokenBody.refresh_token),
displayName: displayName ?? "Google Antigravity",
projectId,
expiresAtMs: expiresIn === null ? null : Date.now() + expiresIn * 1000,
quota,
}),
];
}
export const antigravityDeviceOAuth: OAuth2AddMethod = {
type: "oauth2.0",
export const antigravityAuthorizationCodeOAuth: OAuth2AuthorizationCodeAddMethod = {
type: "oauth2.authorization-code",
id: "google-antigravity",
displayName: {
"en-US": "Sign in with Google (Antigravity)",
"zh-CN": "使用 Google (Antigravity) 登录",
},
description: {
"en-US": "Authorize Antigravity with your Google Account for Gemini & Claude models.",
"zh-CN": "使用 Google 账号完成 Antigravity 授权,畅享 Gemini 与 Claude 模型。",
"en-US": "Authorize Antigravity with your Google Account for Gemini and Claude models.",
"zh-CN": "使用 Google 账号完成 Antigravity 授权,以使用 Gemini 与 Claude 模型。",
},
callback: { port: CALLBACK_PORT, path: CALLBACK_PATH },
begin,
poll,
complete,
};
@@ -0,0 +1,79 @@
import type { PluginContext } from "cursor-byok:plugin";
import { antigravityAuthorizationCodeOAuth } from "./oauth.ts";
function assert(condition: unknown, message: string): asserts condition {
if (!condition) throw new Error(message);
}
function context(requests: Array<{ url: string; body?: string }>): PluginContext {
return {
network: {
fetch: async (url, init = {}) => {
requests.push({ url, body: init.body });
if (url === "https://oauth2.googleapis.com/token") {
return {
status: 200,
headers: {},
body: JSON.stringify({
access_token: "access-token",
refresh_token: "refresh-token",
expires_in: 3600,
}),
};
}
if (url.startsWith("https://www.googleapis.com/oauth2/v1/userinfo")) {
return { status: 200, headers: {}, body: JSON.stringify({ email: "user@example.com" }) };
}
return { status: 500, headers: {}, body: "{}" };
},
stream: () => Promise.reject(new Error("stream is not expected")),
},
signal: new AbortController().signal,
};
}
Deno.test("authorization URL uses Core-owned state, callback, and PKCE challenge", async () => {
const result = await antigravityAuthorizationCodeOAuth.begin(
{
redirectUri: "http://127.0.0.1:51121/oauth-callback",
state: "core-state",
codeChallenge: "core-challenge",
},
context([]),
);
const url = new URL(result.authorizationUrl);
assert(url.searchParams.get("state") === "core-state", "state must come from Core");
assert(
url.searchParams.get("redirect_uri")?.endsWith("/oauth-callback"),
"callback must be forwarded",
);
assert(
url.searchParams.get("code_challenge") === "core-challenge",
"PKCE challenge must be forwarded",
);
assert(url.searchParams.get("code_challenge_method") === "S256", "PKCE must use S256");
});
Deno.test("authorization completion exchanges the code with the Core PKCE verifier", async () => {
const requests: Array<{ url: string; body?: string }> = [];
const resources = await antigravityAuthorizationCodeOAuth.complete(
{ createdAtMs: Date.now() },
{
code: "authorization-code",
redirectUri: "http://127.0.0.1:51121/oauth-callback",
codeVerifier: "core-verifier",
},
context(requests),
);
const tokenRequest = requests.find((request) =>
request.url === "https://oauth2.googleapis.com/token"
);
const body = new URLSearchParams(tokenRequest?.body);
assert(body.get("code") === "authorization-code", "authorization code must be exchanged");
assert(body.get("code_verifier") === "core-verifier", "PKCE verifier must come from Core");
assert(resources.length === 1, "one Google account resource must be returned");
assert(
resources[0].key === "antigravity:user@example.com",
"the account email must be the resource key",
);
});
@@ -2,23 +2,18 @@
"apiVersion": 1,
"id": "dev.cursorbyok.plugins.antigravity-auth",
"name": "Antigravity",
"version": "0.2.0",
"version": "0.3.0",
"author": "Antigravity",
"minAppVersion": "0.1.0",
"icon": "assets/antigravity.svg",
"entry": "main.ts",
"permissions": {
"network": [
"127.0.0.1",
"localhost",
"daily-cloudcode-pa.googleapis.com",
"daily-cloudcode-pa.sandbox.googleapis.com",
"cloudcode-pa.googleapis.com",
"generativelanguage.googleapis.com",
"oauth2.googleapis.com",
"accounts.google.com",
"www.googleapis.com",
"antigravity.google"
"www.googleapis.com"
]
}
}
@@ -88,16 +88,23 @@ function resolveAntigravityModel(modelId: string): string {
}
// 2. Canonical Antigravity-Manager mapping for aliases
if (lower === "claude-3-7-sonnet" || lower === "claude-3-5-sonnet" || lower === "claude-sonnet-4-5") {
if (
lower === "claude-3-7-sonnet" || lower === "claude-3-5-sonnet" || lower === "claude-sonnet-4-5"
) {
return "claude-sonnet-4-6";
}
if (lower === "claude-3-5-haiku" || lower === "claude-haiku-4") {
return "claude-sonnet-4-6";
}
if (lower === "claude-3-7-opus" || lower === "claude-opus-4" || lower === "claude-opus-4.6" || lower === "claude-opus-4-5-thinking") {
if (
lower === "claude-3-7-opus" || lower === "claude-opus-4" || lower === "claude-opus-4.6" ||
lower === "claude-opus-4-5-thinking"
) {
return "claude-opus-4-6-thinking";
}
if (lower === "gpt-4" || lower === "gpt-4o" || lower === "gpt-4o-mini" || lower === "gpt-3.5-turbo") {
if (
lower === "gpt-4" || lower === "gpt-4o" || lower === "gpt-4o-mini" || lower === "gpt-3.5-turbo"
) {
return "gemini-2.5-flash";
}
if (lower === "gemini-2.5-flash-lite") {
@@ -195,7 +202,7 @@ function convertToCloudCodeContents(
messages: LlmMessage[],
): {
contents: Array<{ role: string; parts: Array<Record<string, unknown>> }>;
systemInstruction?: { parts: Array<{ text: string }> };
systemInstruction?: { role: string; parts: Array<{ text: string }> };
} {
const rawContents: Array<{ role: string; parts: Array<Record<string, unknown>> }> = [];
let systemText = instructions || "";
@@ -221,13 +228,17 @@ function convertToCloudCodeContents(
const replayVal = msg.replayState?.providerKind === "antigravity"
? (msg.replayState.value as Record<string, unknown> | null)
: null;
const sig = typeof replayVal?.thoughtSignature === "string" ? replayVal.thoughtSignature : null;
const sig = typeof replayVal?.thoughtSignature === "string"
? replayVal.thoughtSignature
: null;
for (const call of msg.toolCalls) {
parts.push({
functionCall: {
name: call.name,
args: typeof call.arguments === "object" && call.arguments !== null ? call.arguments : {},
args: typeof call.arguments === "object" && call.arguments !== null
? call.arguments
: {},
},
thoughtSignature: sig || "skip_thought_signature_validator",
});
@@ -287,7 +298,9 @@ function convertToCloudCodeContents(
return {
contents,
...(systemText.trim() ? { systemInstruction: { role: "system", parts: [{ text: systemText.trim() }] } } : {}),
...(systemText.trim()
? { systemInstruction: { role: "system", parts: [{ text: systemText.trim() }] } }
: {}),
};
}
@@ -307,20 +320,20 @@ async function streamCloudCode(
const tools = input.request.tools && input.request.tools.length > 0
? [
{
functionDeclarations: input.request.tools.map((t) => ({
name: t.name,
description: t.description || "",
parameters: sanitizeSchema(t.parameters),
})),
},
]
{
functionDeclarations: input.request.tools.map((t) => ({
name: t.name,
description: t.description || "",
parameters: sanitizeSchema(t.parameters),
})),
},
]
: undefined;
const toolConfig = tools
? {
functionCallingConfig: { mode: "AUTO" },
}
functionCallingConfig: { mode: "AUTO" },
}
: undefined;
const payload = {
@@ -348,7 +361,8 @@ async function streamCloudCode(
...ANTIGRAVITY_CLIENT_HEADERS,
};
if (actualModel.toLowerCase().includes("claude")) {
headers["anthropic-beta"] = "claude-code-20250219,interleaved-thinking-2025-05-14,fine-grained-tool-streaming-2025-05-14";
headers["anthropic-beta"] =
"claude-code-20250219,interleaved-thinking-2025-05-14,fine-grained-tool-streaming-2025-05-14";
}
let lastError: Error | null = null;
@@ -370,7 +384,10 @@ async function streamCloudCode(
if (response.status < 200 || response.status >= 300) {
const errorBody = await readBody(response.lines);
lastError = new HttpError(response.status, errorBody);
if (response.status === 503 || response.status === 502 || response.status === 504 || response.status === 404) {
if (
response.status === 503 || response.status === 502 || response.status === 504 ||
response.status === 404
) {
continue;
}
throw lastError;
@@ -382,7 +399,11 @@ async function streamCloudCode(
let hasTools = false;
let toolIndex = 0;
let lastThoughtSignature: string | null = null;
let finalUsage: { inputTokens: number | null; outputTokens: number | null; totalTokens: number | null } | null = null;
let finalUsage: {
inputTokens: number | null;
outputTokens: number | null;
totalTokens: number | null;
} | null = null;
for await (const line of response.lines) {
if (!line.startsWith("data:")) continue;
@@ -403,7 +424,9 @@ async function streamCloudCode(
if (usage) {
finalUsage = {
inputTokens: typeof usage.promptTokenCount === "number" ? usage.promptTokenCount : null,
outputTokens: typeof usage.candidatesTokenCount === "number" ? usage.candidatesTokenCount : null,
outputTokens: typeof usage.candidatesTokenCount === "number"
? usage.candidatesTokenCount
: null,
totalTokens: typeof usage.totalTokenCount === "number" ? usage.totalTokenCount : null,
};
}
@@ -485,7 +508,9 @@ async function streamCloudCode(
}
}
const finishReason = typeof candidate?.finishReason === "string" ? candidate.finishReason : null;
const finishReason = typeof candidate?.finishReason === "string"
? candidate.finishReason
: null;
if (finishReason) {
if (thinkingStarted) {
thinkingStarted = false;
@@ -517,7 +542,8 @@ async function streamCloudCode(
});
finalUsage = null;
}
const isTool = finishReason === "STOP" && (hasTools || parts?.some((p) => p.functionCall));
const isTool = finishReason === "STOP" &&
(hasTools || parts?.some((p) => p.functionCall));
output.emit({
type: "done",
reason: isTool ? "tool-use" : "stop",
@@ -610,17 +636,30 @@ async function invoke(
try {
await streamCloudCode(data.accessToken, projectId, input.model.id, input, output, context);
return patchData
? { status: "completed", patch: { privateData: patchData as unknown as JsonValue, state: { status: "ready" } } }
? {
status: "completed",
patch: { privateData: patchData as unknown as JsonValue, state: { status: "ready" } },
}
: { status: "completed" };
} catch (error) {
if (error instanceof HttpError) {
if ((error.status === 401 || error.status === 403) && data.refreshToken && !isQuotaHttpError(error)) {
if (
(error.status === 401 || error.status === 403) && data.refreshToken &&
!isQuotaHttpError(error)
) {
try {
const refreshed = await refreshAccount(input.resource, context);
if (refreshed.privateData) {
const freshData = refreshed.privateData as unknown as AccountData;
const freshProj = freshData.projectId ?? projectId;
await streamCloudCode(freshData.accessToken, freshProj, input.model.id, input, output, context);
await streamCloudCode(
freshData.accessToken,
freshProj,
input.model.id,
input,
output,
context,
);
return {
status: "completed",
patch: { privateData: freshData as unknown as JsonValue, state: { status: "ready" } },
@@ -1,4 +1,4 @@
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
import type { JsonValue, NetworkRequestInit, PluginContext } from "cursor-byok:plugin";
import type {
ResourceDraft,
ResourceImportFile,
@@ -15,11 +15,25 @@ import {
ANTIGRAVITY_ENDPOINTS,
ANTIGRAVITY_USER_AGENT,
} from "./models.ts";
import { CLIENT_ID, CLIENT_SECRET } from "./google_oauth.ts";
export const RESOURCE_TYPE = "antigravity-account";
const REFRESH_TOKEN_URL = "https://oauth2.googleapis.com/token";
async function fetchText(
network: PluginContext["network"] | undefined,
url: string,
init: NetworkRequestInit,
): Promise<{ status: number; body: string }> {
if (network) {
const response = await network.fetch(url, init);
return { status: response.status, body: response.body };
}
const response = await fetch(url, init);
return { status: response.status, body: await response.text() };
}
export type QuotaMetric = {
remainingPercent: number;
resetAtMs: number | null;
@@ -73,12 +87,15 @@ export async function fetchAccountProjectAndTier(
const project = text(body?.cloudaicompanionProject);
const paid = object(body?.paidTier);
const current = object(body?.currentTier);
const tierName = text(paid?.name) ?? text(paid?.id) ?? text(current?.name) ?? text(current?.id);
const tierName = text(paid?.name) ?? text(paid?.id) ?? text(current?.name) ??
text(current?.id);
let planLabel = "FREE";
if (tierName) {
const lower = tierName.toLowerCase();
if (lower.includes("ultra")) planLabel = "ULTRA";
else if (lower.includes("pro") || lower.includes("premium") || lower.includes("advanced")) planLabel = "PRO";
else if (
lower.includes("pro") || lower.includes("premium") || lower.includes("advanced")
) planLabel = "PRO";
}
return { projectId: project ?? "bamboo-precept-lgxtn", planLabel };
}
@@ -121,7 +138,9 @@ export async function queryAccountQuota(
for (const [key, value] of Object.entries(models)) {
const info = object(value);
const quota = object(info?.quotaInfo);
const fraction = typeof quota?.remainingFraction === "number" ? quota.remainingFraction : null;
const fraction = typeof quota?.remainingFraction === "number"
? quota.remainingFraction
: null;
const resetTime = text(quota?.resetTime);
const resetAtMs = resetTime ? Date.parse(resetTime) : null;
if (fraction === null) continue;
@@ -147,8 +166,12 @@ export async function queryAccountQuota(
limitReached: false,
coolingUntilMs: null,
updatedAtMs: Date.now(),
claude: claudeFraction !== null ? { remainingPercent: Math.round(claudeFraction * 100), resetAtMs: claudeResetAtMs } : null,
gemini: geminiFraction !== null ? { remainingPercent: Math.round(geminiFraction * 100), resetAtMs: geminiResetAtMs } : null,
claude: claudeFraction !== null
? { remainingPercent: Math.round(claudeFraction * 100), resetAtMs: claudeResetAtMs }
: null,
gemini: geminiFraction !== null
? { remainingPercent: Math.round(geminiFraction * 100), resetAtMs: geminiResetAtMs }
: null,
},
};
} catch {
@@ -235,7 +258,8 @@ export async function accountIdentity(
const identity = (providedDisplayName && !providedDisplayName.includes("Antigravity"))
? providedDisplayName
: (email ?? sub ?? fingerprint);
const displayName = providedDisplayName ?? email ?? name ?? (token.startsWith("AIza") ? `API Key (${fingerprint.slice(0, 6)})` : identity);
const displayName = providedDisplayName ?? email ?? name ??
(token.startsWith("AIza") ? `API Key (${fingerprint.slice(0, 6)})` : identity);
return { key: `antigravity:${identity}`, displayName };
}
@@ -246,7 +270,8 @@ export async function credentialDraft(credential: CredentialCandidate): Promise<
refreshToken: credential.refreshToken,
displayName: credential.displayName ?? identity.displayName,
projectId: credential.projectId ?? "bamboo-precept-lgxtn",
expiresAtMs: credential.expiresAtMs ?? (credential.refreshToken ? Date.now() + 3500 * 1000 : null),
expiresAtMs: credential.expiresAtMs ??
(credential.refreshToken ? Date.now() + 3500 * 1000 : null),
quota: credential.quota ?? null,
};
return { key: identity.key, privateData: data as unknown as JsonValue };
@@ -346,8 +371,6 @@ export function presentAccount(resource: ResourceSnapshot): ResourceView {
};
}
import { CLIENT_ID, CLIENT_SECRET } from "./oauth.ts";
export async function refreshAccount(
resource: ResourceSnapshot,
context: PluginContext,
@@ -378,7 +401,10 @@ export async function refreshAccount(
// Only mark invalid if token is revoked or client is invalid
if (bodyText.includes("invalid_grant") || bodyText.includes("unauthorized_client")) {
return {
state: { status: "invalid", message: "Google authorization expired or revoked; please sign in again" },
state: {
status: "invalid",
message: "Google authorization expired or revoked; please sign in again",
},
};
}
// On network glitches or temporary Google server errors, keep ready
@@ -461,18 +487,18 @@ function collectCredentials(value: unknown, output: CredentialCandidate[]): void
firstText(item, ["refresh", "refresh_token", "refreshToken"]);
const displayName = firstText(item, ["email", "display_name", "displayName", "name"]) ??
firstText(tokens, ["email", "display_name", "displayName", "name"]);
const projectId = firstText(item, ["project", "projectId", "project_id", "cloudaicompanionProject"]) ??
firstText(tokens, ["project", "projectId", "project_id", "cloudaicompanionProject"]);
const projectId =
firstText(item, ["project", "projectId", "project_id", "cloudaicompanionProject"]) ??
firstText(tokens, ["project", "projectId", "project_id", "cloudaicompanionProject"]);
if (!accessToken && !refreshToken) return;
if (!accessToken && refreshToken) {
accessToken = refreshToken;
}
if (!accessToken) return;
output.push({ accessToken, refreshToken, displayName, projectId });
}
import { REDIRECT_URI } from "./oauth.ts";
export async function parseCredentialFiles(
files: ResourceImportFile[],
network?: PluginContext["network"],
@@ -486,41 +512,6 @@ export async function parseCredentialFiles(
const raw = file.content.trim();
if (!raw) continue;
// Check if user pasted the Google OAuth callback URL or raw auth code (4/0ATs...)
const authCodeMatch = raw.match(/code=([40][a-zA-Z0-9_\-%]+)/) || (raw.startsWith("4/") ? [null, raw] : null);
if (authCodeMatch?.[1]) {
const code = decodeURIComponent(authCodeMatch[1]);
try {
const fetcher = network ? (url: string, init: RequestInit) => network.fetch(url, init) : (url: string, init: RequestInit) => fetch(url, init).then(async (r) => ({ status: r.status, body: await r.text() }));
const response = await fetcher(REFRESH_TOKEN_URL, {
method: "POST",
headers: {
accept: "application/json",
"content-type": "application/x-www-form-urlencoded",
},
body: new URLSearchParams({
client_id: CLIENT_ID,
client_secret: CLIENT_SECRET,
code,
grant_type: "authorization_code",
redirect_uri: REDIRECT_URI,
}).toString(),
});
const tokenBody = object(JSON.parse(response.body));
if (response.status >= 200 && response.status < 300 && text(tokenBody?.access_token)) {
credentials.push({
accessToken: text(tokenBody?.access_token)!,
refreshToken: text(tokenBody?.refresh_token),
displayName: text(tokenBody?.email) ?? "Google Antigravity Account",
expiresAtMs: Date.now() + ((typeof tokenBody?.expires_in === "number" ? tokenBody.expires_in : 3600) * 1000),
});
continue;
}
} catch {
// Fallback to normal parsing
}
}
// Check if file is raw API key or JWT token string
if (raw.startsWith("AIza") || (raw.split(".").length === 3 && !raw.includes(" "))) {
credentials.push({ accessToken: raw, refreshToken: null, displayName: file.name });
@@ -532,9 +523,15 @@ export async function parseCredentialFiles(
try {
content = JSON.parse(raw);
} catch {
const envMatch = raw.match(/(?:API_KEY|TOKEN|GEMINI_API_KEY|GOOGLE_API_KEY|ANTIGRAVITY_API_KEY)\s*=\s*["']?([^"'\r\n]+)/i);
const envMatch = raw.match(
/(?:API_KEY|TOKEN|GEMINI_API_KEY|GOOGLE_API_KEY|ANTIGRAVITY_API_KEY)\s*=\s*["']?([^"'\r\n]+)/i,
);
if (envMatch?.[1]) {
credentials.push({ accessToken: envMatch[1].trim(), refreshToken: null, displayName: file.name });
credentials.push({
accessToken: envMatch[1].trim(),
refreshToken: null,
displayName: file.name,
});
continue;
}
const keyMatch = raw.match(/AIza[0-9A-Za-z-_]{35}/);
@@ -560,11 +557,7 @@ export async function parseCredentialFiles(
for (const candidate of found) {
if (candidate.refreshToken && candidate.accessToken === candidate.refreshToken) {
try {
const fetcher = network
? (url: string, init: RequestInit) => network.fetch(url, init)
: (url: string, init: RequestInit) =>
fetch(url, init).then(async (r) => ({ status: r.status, body: await r.text() }));
const response = await fetcher(REFRESH_TOKEN_URL, {
const response = await fetchText(network, REFRESH_TOKEN_URL, {
method: "POST",
headers: {
accept: "application/json",
@@ -581,7 +574,8 @@ export async function parseCredentialFiles(
if (response.status >= 200 && response.status < 300 && text(body?.access_token)) {
candidate.accessToken = text(body?.access_token)!;
candidate.refreshToken = text(body?.refresh_token) ?? candidate.refreshToken;
candidate.expiresAtMs = Date.now() + ((typeof body?.expires_in === "number" ? body.expires_in : 3600) * 1000);
candidate.expiresAtMs = Date.now() +
((typeof body?.expires_in === "number" ? body.expires_in : 3600) * 1000);
}
} catch {
// Keep placeholder
@@ -599,15 +593,22 @@ export const credentialImport: ResourceImportSupport = {
"zh-CN": "导入 Google / Antigravity 凭证",
},
description: {
"en-US": "Import a JSON, TXT, or Callback URL containing Antigravity tokens or Google API keys.",
"zh-CN": "导入包含 Antigravity Token、Google API Key 或授权回调 URL 的 JSON/TXT 文件。",
"en-US":
"Import a JSON, TXT, or environment file containing Antigravity tokens or Google API keys.",
"zh-CN": "导入包含 Antigravity Token 或 Google API Key 的 JSON、TXT 或环境变量文件。",
},
accept: [".json", ".txt", ".key", ".env"],
multiple: true,
parse: async (files: ResourceImportFile[], context: PluginContext): Promise<ResourceImportResult> => {
parse: async (
files: ResourceImportFile[],
context: PluginContext,
): Promise<ResourceImportResult> => {
const { credentials, warnings } = await parseCredentialFiles(files, context.network);
if (credentials.length === 0) {
throw new Error(warnings.join("; ") || "credential file does not contain a valid token, authorization code, or API key");
throw new Error(
warnings.join("; ") ||
"credential file does not contain a valid token or API key",
);
}
const drafts = await Promise.all(
credentials.map(async (c) => {
@@ -8,6 +8,7 @@ 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 { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
import { codexProvider, isQuotaError } from "./provider.ts";
import {
accountIdentity,
@@ -305,6 +306,43 @@ Deno.test("invoke streams normalized events from the Codex Responses API", async
]);
});
Deno.test("reasoning replay projects response items to valid input items", () => {
const replayRequest = request();
replayRequest.messages = [{
role: "assistant",
text: "",
thinking: "",
replayState: {
providerKind: "openai_responses",
value: {
items: [{
type: "reasoning",
id: "item-1",
status: "completed",
summary: [{ type: "summary_text", text: "why" }],
content: [],
encrypted_content: "opaque",
output_only: true,
}],
},
},
toolCalls: [],
}];
const body = buildResponsesBody({
url: "https://example.com/responses",
model: "gpt-test",
request: replayRequest,
});
assertEquals(body.input, [{
type: "reasoning",
id: "item-1",
summary: [{ type: "summary_text", text: "why" }],
content: [],
encrypted_content: "opaque",
}]);
});
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({
@@ -277,7 +277,9 @@ Deno.test("invoke streams normalized events from the xAI Chat Completions API",
);
assertEquals(result, { status: "completed" });
const body = JSON.parse(requestBody) as Record<string, unknown>;
assert(!("prompt_cache_key" in body), "standard OpenAI chat completion does not include prompt_cache_key");
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}`);
+67
View File
@@ -0,0 +1,67 @@
# Git Commit Message Generation Guide
## Role and objective
You are a Git commit message generator. Given a Git diff, output only the commit message itself, without explanations, preambles, quotation marks, or additional text. Your entire response will be passed directly to `git commit`.
## General rules for subjects
- Use the present tense and describe the key change in the diff precisely.
- Focus on what changed instead of listing file names.
- Be specific: include concrete details such as package names, versions, or features, and avoid vague descriptions.
- Exclude unnecessary content such as translation notes.
- Keep the subject at or below 50 characters.
- Write the commit message in English regardless of the language used in the diff.
- Output only the commit message text, without quotation marks, formatting wrappers, explanations, or preambles.
## Output format
Choose exactly one format based on `type`:
| type | Format template |
| ----------------- | ----------------------------------------------------------- |
| plain | `<commit message>` |
| conventional | `<type>[optional (<scope>)]: <commit message>` |
| conventional+body | `<type>[optional (<scope>)]: <commit message subject>` |
| gitmoji | `:emoji: <commit message>` |
| subject+body | `<commit message subject>` |
For `conventional` and `conventional+body`, the subject must begin with a lowercase letter. The output must strictly follow the selected format.
## Conventional type selection
Choose the single type that best matches the diff. The type must be lowercase, such as `feat`, never `Feat` or `FEAT`.
```json
{
"docs": "documentation-only changes",
"style": "changes that do not affect code meaning, such as whitespace, formatting, or missing semicolons",
"refactor": "code structure improvements that do not change behavior, such as renaming, restructuring methods, or extracting functions",
"perf": "code changes that improve performance",
"test": "adding missing tests or correcting existing tests",
"build": "changes that affect the build system or external dependencies",
"ci": "changes to CI configuration and scripts",
"chore": "other changes that do not modify src or test files",
"revert": "reverting a previous commit",
"feat": "a new feature",
"fix": "a bug fix"
}
```
- For `conventional`, output the complete conventional subject line.
- For `conventional+body`, output only the conventional subject line; the body is generated separately.
## Body generation rules
When a commit subject is already provided and a description is requested, output only the commit body:
- Keep it concise: use 3–6 short bullet points, one per line, or 2–4 short sentences.
- Use the present tense and focus on what changed and why.
- Keep every line at or below 72 characters. Indent wrapped bullet lines by two spaces so they align with the bullet text.
- Do not repeat the subject or add meta commentary such as “This commit”.
- Write in English.
- Output only the body, without any additional text.
- Describe concrete changes clearly; avoid vague phrases such as “update functionality” or “modify resources”.
- Every commit subject must have a prefix and must not use emoji. If the changes cover separate concerns, such as visual improvements and bug fixes, split them into separate entries, for example:
- `fix(<specific area>): fix the xxx issue`
- `chore(<specific area>): update visual assets`
+65
View File
@@ -0,0 +1,65 @@
# Git 提交信息生成指南
## 角色与目标
你是一个 git 提交信息生成器。给你一段 git diff,你只输出提交信息本身,不要有任何解释、前言、引号或额外文本。你的整个回复会被直接传给 `git commit`。
## 通用规则(生成标题时)
- 使用现在时(present tense),精准描述本次 diff 的关键改动。
- 关注「改了什么」,而不是罗列文件名。
- 要具体:包含具体细节(包名、版本、功能点),避免笼统表述。
- 排除任何不必要的内容(如翻译说明)。
- 标题最长不超过 50 个字符。
- 提交信息语言:简体中文(无论 diff 中的内容是什么语言,输出一律用简体中文)。
- 只输出提交信息文本本身,不要用引号或其他格式包裹,不要加解释或前言。
## 输出格式(按 type 选择其一)
| type | 格式模板 |
| ------------------ | ------------------------------------------------------------ |
| plain | `<commit message>` |
| conventional | `<type>[optional (<scope>)]: <commit message>`(主题须以小写字母开头) |
| conventional+body | `<type>[optional (<scope>)]: <commit message subject>`(主题须以小写字母开头,body 单独生成) |
| gitmoji | `:emoji: <commit message>` |
| subject+body | `<commit message subject>`(body 单独生成) |
输出必须严格符合所选 type 对应的格式。
## Conventional 类型选择
从下面的「类型-描述」中选择一个最贴合本次 diff 的类型。重要:类型必须全小写(例如 `feat`,而不是 `Feat` 或 `FEAT`)。
```json
{
"docs": "仅文档变更",
"style": "不影响代码含义的变更(空白、格式、缺失分号等)",
"refactor": "改善代码结构但不改变功能的变更(重命名、重构类/方法、抽取函数等)",
"perf": "提升性能的代码变更",
"test": "新增缺失的测试或修正已有测试",
"build": "影响构建系统或外部依赖的变更",
"ci": "对 CI 配置文件和脚本的变更",
"chore": "不修改 src 或 test 文件的其他变更",
"revert": "回退某个之前的提交",
"feat": "新功能",
"fix": "缺陷修复"
}
```
- `conventional`:直接按上表选择类型并输出完整主题行。
- `conventional+body`:只输出 conventional 主题行,body 会单独生成。
## 描述(body)生成规则
当已有提交标题、需要生成描述时,给你标题与 diff,你只输出提交描述正文:
- 简洁:使用 3–6 条要点(每条一行短句),或 2–4 句短句,不要长段落。
- 用现在时聚焦「改了什么、为什么」。
- 每行最多 72 个字符;当某条要点换行时,续行缩进 2 个空格,与要点文字对齐。
- 不要重复标题,不要元评论(如「本次提交……」)。
- 语言:简体中文。
- 只输出提交描述正文,不要有其他内容。
- 修改点要列举清楚明白,不能笼统地说“更新功能、修改资源”。
- 提交必须有前缀,且不允许使用 emoji。如果存在多个提交功能,例如美化和缺陷修复,应拆分成多条,例如:
- `fix(<具体修改项>): 修复 xxx 问题`
- `chore(<具体修改项>): 修改美术资源`
+5 -1
View File
@@ -184,7 +184,11 @@ pub async fn append(
request: DecodedAppend,
parent: Option<TransportParent>,
) -> Result<ai::BidiAppendResponse> {
let handle = registry.get_or_create(&request.request_id).await?;
let replace_closing = request.model_id().is_some();
let handle = registry
.get_or_create_for_append(&request.request_id, replace_closing)
.await?;
let _admission = handle.admit()?;
if let Some(conversation_id) = request.conversation_id() {
handle.set_conversation_id(conversation_id)?;
}
+162 -26
View File
@@ -20,15 +20,19 @@ use crate::{
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
account, analytics, commit_message, compatibility, entitlement::FreeEntitlementCache,
knowledge, model_catalog, server_config, tab,
},
transport::{TransportParent, TransportRegistry},
},
Result,
};
pub fn router(registry: TransportRegistry) -> Result<Router> {
let proxy = CursorProxy::cursor(registry.store().clone())?;
pub fn router(
registry: TransportRegistry,
clients: crate::network::NetworkClients,
) -> Result<Router> {
let proxy = CursorProxy::cursor(clients);
let knowledge = knowledge::KnowledgeService::managed()?;
Ok(router_with_proxy(registry, proxy, knowledge))
}
@@ -39,10 +43,35 @@ fn router_with_proxy(
knowledge_service: knowledge::KnowledgeService,
) -> Router {
let web_cache = registry.web_cache().router();
let free_entitlements = FreeEntitlementCache::default();
Router::new()
.route("/__byok-api__/healthz", get(health))
.route("/agent.v1.AgentService/RunSSE", post(run_sse_handler))
.route("/aiserver.v1.BidiService/BidiAppend", post(bidi_handler))
.route(
"/aiserver.v1.AiService/AvailableDocs",
post(compatibility::available_docs),
)
.route(
"/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
post(compatibility::effective_user_plugins),
)
.route(
"/aiserver.v1.DashboardService/GetUserPrivacyMode",
post(compatibility::user_privacy_mode),
)
.route(
"/agent.v1.AgentService/UpdateConversationMetadata",
post(compatibility::update_conversation_metadata),
)
.route(
"/aiserver.v1.AiService/GetServerConfig",
post(server_config::get),
)
.route(
"/aiserver.v1.ServerConfigService/GetServerConfig",
post(server_config::get),
)
.route(
"/aiserver.v1.AiService/AvailableModels",
post(model_catalog::available_models),
@@ -55,10 +84,38 @@ fn router_with_proxy(
"/aiserver.v1.AiService/GetUsableModels",
post(model_catalog::usable_models),
)
.route(
"/aiserver.v1.AiService/WriteGitCommitMessage",
post(commit_message::write_git_commit_message),
)
.route(
"/aiserver.v1.NetworkService/IsConnected",
post(is_connected),
)
.route(
"/agent.v1.AgentService/GetDefaultModelForCli",
post(model_catalog::default_model_for_cli),
)
.route(
"/aiserver.v1.AiService/GetDefaultModelForCli",
post(model_catalog::default_model_for_cli),
)
.route(
"/aiserver.v1.AiService/GetDefaultModel",
post(model_catalog::default_model),
)
.route(
"/aiserver.v1.AiService/GetDefaultModelNudgeData",
post(model_catalog::default_model_nudge),
)
.route(
"/aiserver.v1.AuthService/GetEmail",
post(account::get_email),
)
.route(
"/aiserver.v1.AuthService/GetUserMeta",
post(account::get_user_meta),
)
.route("/aiserver.v1.DashboardService/GetMe", post(account::get_me))
.route(
"/aiserver.v1.DashboardService/GetTeams",
@@ -97,6 +154,7 @@ fn router_with_proxy(
post(analytics::bootstrap_statsig),
)
.route("/auth/full_stripe_profile", get(account::stripe_profile))
.route("/auth/stripe_profile", get(account::stripe_profile))
.merge(tab::router())
.route_layer(DefaultBodyLimit::disable())
.route_layer(RequestDecompressionLayer::new())
@@ -104,6 +162,7 @@ fn router_with_proxy(
.method_not_allowed_fallback(proxy::forward)
.layer(Extension(proxy))
.layer(Extension(knowledge_service))
.layer(Extension(free_entitlements))
.with_state(registry)
.merge(web_cache)
}
@@ -112,6 +171,23 @@ async fn health() -> StatusCode {
StatusCode::NO_CONTENT
}
/// `NetworkService/IsConnected` probe. Cursor's always-local extension checks
/// connectivity roughly 10s after any slow request starts; a 404/error here is
/// treated as "network disconnected" and aborts in-flight work (e.g. commit
/// message generation) even while the model is still streaming. Always answer
/// connected with an empty `IsConnectedResponse` so local BYOK generation is
/// never cancelled by this probe.
async fn is_connected() -> Result<Response<Body>> {
let payload = connect::encode_message(&ai::IsConnectedResponse {})?;
let mut response = Response::new(Body::from(payload));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
async fn run_sse_handler(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
@@ -120,16 +196,13 @@ async fn run_sse_handler(
let (parts, body) = buffered(request).await?;
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
let route = registry.wait_route(&request.request_id).await;
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
if let Some(trace) = &trace {
trace
.request(
"run_sse_request",
&body,
serde_json::json!({"request_id": request.request_id}),
)
.await;
}
let trace = registry.trace(&request.request_id);
trace.resume();
trace.request(
"run_sse_request",
body.clone(),
serde_json::json!({"request_id": request.request_id}),
);
match route {
crate::cursor::transport::TransportRoute::Local => {
run_sse::stream(&registry, &request.request_id).await
@@ -140,7 +213,14 @@ async fn run_sse_handler(
Request::from_parts(parts, Body::from(body)),
)
.await?;
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
Ok(run_sse::upstream(
registry,
request.request_id,
generation,
response,
Some(trace),
)
.await)
}
}
}
@@ -156,6 +236,7 @@ async fn bidi_handler(
let first_model = decoded.model_id().map(str::to_owned);
let conversation_id = decoded.conversation_id().map(str::to_owned);
let trace_metadata = decoded.trace_metadata();
let trace = registry.trace(&decoded.request_id);
let local = if let Some(model_id) = decoded.model_id() {
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
@@ -180,14 +261,18 @@ async fn bidi_handler(
} else if registry.upstream(&decoded.request_id).await {
false
} else {
trace.resume();
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, false, "missing_transport", None),
);
return Err(crate::Error::Protocol(
"first BidiAppend message must select a model".into(),
));
};
let trace = if first_model.is_some() {
CursorTraceRecorder::begin(
registry.store().clone(),
&decoded.request_id,
if first_model.is_some() {
trace.begin(
conversation_id.as_deref(),
if local {
"local_byok"
@@ -195,26 +280,61 @@ async fn bidi_handler(
"cursor_official"
},
first_model.as_deref(),
)
.await
);
} else {
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
};
if let Some(trace) = &trace {
trace.request("bidi_request", &body, trace_metadata).await;
trace.resume();
}
if !local {
if first_model.is_some() {
registry.mark_upstream(&decoded.request_id).await;
}
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, true, "upstream", None),
);
return proxy::forward(
Extension(proxy),
Request::from_parts(parts, Body::from(body)),
)
.await;
}
let parent = parent_headers(&parts.headers)?;
bidi::append(&registry, decoded, parent).await?;
let parent = match parent_headers(&parts.headers) {
Ok(parent) => parent,
Err(error) => {
trace.request(
"bidi_request",
body,
trace_outcome(
trace_metadata,
false,
"invalid_parent",
Some(error.to_string()),
),
);
return Err(error);
}
};
match bidi::append(&registry, decoded, parent).await {
Ok(_) => trace.request(
"bidi_request",
body,
trace_outcome(trace_metadata, true, "local", None),
),
Err(error) => {
trace.request(
"bidi_request",
body,
trace_outcome(
trace_metadata,
false,
"command_rejected",
Some(error.to_string()),
),
);
return Err(error);
}
}
let mut response = Response::new(axum::body::Body::empty());
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
@@ -224,6 +344,22 @@ async fn bidi_handler(
Ok(response)
}
fn trace_outcome(
mut metadata: serde_json::Value,
accepted: bool,
route_outcome: &str,
error: Option<String>,
) -> serde_json::Value {
if let Some(metadata) = metadata.as_object_mut() {
metadata.insert("accepted".into(), accepted.into());
metadata.insert("route_outcome".into(), route_outcome.into());
if let Some(error) = error {
metadata.insert("error".into(), error.into());
}
}
metadata
}
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)
+6 -15
View File
@@ -14,8 +14,7 @@ pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url";
#[derive(Clone)]
pub struct CursorProxy {
client: Option<reqwest::Client>,
store: Option<crate::store::Store>,
clients: crate::network::NetworkClients,
upstream: String,
}
@@ -47,23 +46,15 @@ impl BufferedResponse {
}
impl CursorProxy {
pub fn cursor(store: crate::store::Store) -> Result<Self> {
Ok(Self {
client: None,
store: Some(store),
pub fn cursor(clients: crate::network::NetworkClients) -> Self {
Self {
clients,
upstream: CURSOR_UPSTREAM.into(),
})
}
}
async fn client(&self) -> Result<reqwest::Client> {
match (&self.client, &self.store) {
(Some(client), _) => Ok(client.clone()),
(_, Some(store)) => Ok(crate::network::client_builder(store)
.await?
.redirect(reqwest::redirect::Policy::none())
.build()?),
_ => unreachable!("Cursor proxy always has a client or store"),
}
self.clients.cursor_client().await
}
}
+5 -5
View File
@@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Re
let receiver = handle.subscribe();
let trace = handle.trace().cloned();
if let Some(trace) = &trace {
trace.response_started(StatusCode::OK.as_u16()).await;
trace.response_started(StatusCode::OK.as_u16());
}
let body_stream = local_body_stream(receiver, handle, trace);
let mut response = Response::new(Body::from_stream(body_stream));
@@ -133,7 +133,7 @@ pub async fn upstream(
) -> Response<Body> {
let (parts, body) = response.into_parts();
if let Some(trace) = &trace {
trace.response_started(parts.status.as_u16()).await;
trace.response_started(parts.status.as_u16());
}
let stream = async_stream::stream! {
let _guard = UpstreamRunGuard {
@@ -180,15 +180,15 @@ impl TraceStreamSink {
while let Some(event) = receiver.recv().await {
match event {
TraceStreamEvent::Chunk(chunk) => {
trace.response_chunk(source, &chunk).await;
trace.response_chunk(source, chunk);
}
TraceStreamEvent::Finish(error) => {
trace.finish(error.as_deref()).await;
trace.finish(error.as_deref());
return;
}
}
}
trace.finish(None).await;
trace.finish(None);
});
Self {
sender: Some(sender),
+3 -3
View File
@@ -1,7 +1,7 @@
//! Builds the top-level server router.
use crate::{cursor::transport::TransportRegistry, Result};
use crate::{cursor::transport::TransportRegistry, network::NetworkClients, Result};
pub fn router(registry: TransportRegistry) -> Result<axum::Router> {
super::cursor::router(registry)
pub fn router(registry: TransportRegistry, clients: NetworkClients) -> Result<axum::Router> {
super::cursor::router(registry, clients)
}
+11 -4
View File
@@ -44,9 +44,11 @@ impl App {
plugin_runtime.clone(),
config.app_version.clone(),
)?;
let clients = crate::network::NetworkClients::new(store.clone());
let provider = std::sync::Arc::new(ProviderRouter::new(
store.clone(),
plugins.clone(),
clients.clone(),
config.provider_request_timeout,
config.provider_stream_idle_timeout,
));
@@ -58,10 +60,16 @@ impl App {
plugins.clone(),
crate::config::managed_data_dir()?.join("rules"),
);
let control =
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
let control = control::ControlService::new(
store.clone(),
provider,
plugin_runtime,
plugins,
clients.clone(),
config.app_version.clone(),
)?;
let harness = control.cursor_harness().clone();
let mut router = api::router(registry.clone())?;
let mut router = api::router(registry.clone(), clients)?;
router = match &config.console {
Some(ConsoleSource::Directory(directory)) => {
router.merge(control::web_router(control.clone(), directory))
@@ -125,7 +133,6 @@ impl App {
self.registry.web_cache().set_service_addr(address);
self.harness.set_backend_addr(address);
tracing::info!(%address, "cursor server listening");
tokio::spawn(crate::plugin::oauth_callback::start_oauth_callback_server(shutdown.clone()));
let registry = self.registry;
let harness = self.harness;
let graceful = shutdown.clone();
+262 -2
View File
@@ -1,15 +1,25 @@
//! Implements advertisement configuration endpoints.
//! Advertisement service contract and desktop HTTP handler.
use std::{
collections::{BTreeSet, HashMap},
path::{Path as FilePath, PathBuf},
};
use axum::{
body::Body,
extract::{Path, State},
http::{HeaderMap, StatusCode},
http::{header, HeaderMap, HeaderValue, Response, StatusCode},
Json,
};
use bytes::BytesMut;
use futures_util::StreamExt;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use url::Url;
use uuid::Uuid;
use crate::{Error, Result};
use crate::{config, Error, Result};
use super::ControlService;
@@ -17,11 +27,15 @@ use super::ControlService;
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
// pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu";
pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
pub(super) const DISABLED_AD_IDS_HEADER: &str = "disable-ad-ids";
pub(super) const LANGUAGE_HEADER: &str = "accept-language";
const ADS_IMAGE_ROUTE: &str = "/__byok-api__/api/ads/images";
const MAX_AD_IMAGE_BYTES: usize = 10 * 1024 * 1024;
const IMAGE_EXTENSIONS: &[&str] = &["png", "jpg", "gif", "webp", "avif"];
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdRuntime {
@@ -103,6 +117,137 @@ impl AdRuntime {
}
Ok(self)
}
pub(super) async fn cache_images(&mut self, client: &reqwest::Client) {
let cache_dir = match config::managed_data_dir() {
Ok(path) => path.join("ads"),
Err(error) => {
tracing::warn!(%error, "failed to resolve advertisement image cache directory");
return;
}
};
let urls = self
.slots
.iter()
.flat_map(|slot| [&slot.target.image_url, &slot.content.image_url])
.cloned()
.collect::<BTreeSet<_>>();
let downloads = futures_util::future::join_all(urls.iter().map(|url| {
let cache_dir = &cache_dir;
async move {
let result = cache_image(client, cache_dir, url).await;
(url, result)
}
}))
.await;
let mut cached_urls = HashMap::new();
for (url, result) in downloads {
match result {
Ok(cached_url) => {
cached_urls.insert(url.as_str(), cached_url);
}
Err(error) => {
tracing::warn!(%error, image_url = %url, "failed to cache advertisement image")
}
}
}
for slot in &mut self.slots {
if let Some(url) = cached_urls.get(slot.target.image_url.as_str()) {
slot.target.image_url.clone_from(url);
}
if let Some(url) = cached_urls.get(slot.content.image_url.as_str()) {
slot.content.image_url.clone_from(url);
}
}
}
}
async fn cache_image(client: &reqwest::Client, cache_dir: &FilePath, url: &str) -> Result<String> {
tokio::fs::create_dir_all(cache_dir).await?;
let hash = hex::encode(Sha256::digest(url.as_bytes()));
if let Some(file_name) = cached_file_name(cache_dir, &hash).await {
return Ok(format!("{ADS_IMAGE_ROUTE}/{file_name}"));
}
let response = client
.get(url)
.timeout(std::time::Duration::from_secs(10))
.send()
.await?;
if !response.status().is_success() {
return Err(Error::Provider(format!(
"advertisement image download failed ({})",
response.status()
)));
}
let content_type = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.and_then(image_extension)
.ok_or_else(|| {
Error::Provider("advertisement image has an unsupported content type".into())
})?;
if response
.content_length()
.is_some_and(|length| length > MAX_AD_IMAGE_BYTES as u64)
{
return Err(Error::Provider("advertisement image exceeds 10 MiB".into()));
}
let mut bytes = BytesMut::new();
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if bytes.len() + chunk.len() > MAX_AD_IMAGE_BYTES {
return Err(Error::Provider("advertisement image exceeds 10 MiB".into()));
}
bytes.extend_from_slice(&chunk);
}
let file_name = format!("{hash}.{content_type}");
let destination = cache_dir.join(&file_name);
let temporary = cache_dir.join(format!(".{file_name}.{}.tmp", Uuid::new_v4()));
tokio::fs::write(&temporary, &bytes).await?;
if let Err(error) = tokio::fs::rename(&temporary, &destination).await {
if !destination.exists() {
let _ = tokio::fs::remove_file(&temporary).await;
return Err(error.into());
}
let _ = tokio::fs::remove_file(&temporary).await;
}
Ok(format!("{ADS_IMAGE_ROUTE}/{file_name}"))
}
async fn cached_file_name(cache_dir: &FilePath, hash: &str) -> Option<String> {
for extension in IMAGE_EXTENSIONS {
let file_name = format!("{hash}.{extension}");
if tokio::fs::metadata(cache_dir.join(&file_name))
.await
.is_ok()
{
return Some(file_name);
}
}
None
}
fn image_extension(content_type: &str) -> Option<&'static str> {
match content_type
.split(';')
.next()
.unwrap_or_default()
.trim()
.to_ascii_lowercase()
.as_str()
{
"image/png" => Some("png"),
"image/jpeg" => Some("jpg"),
"image/gif" => Some("gif"),
"image/webp" => Some("webp"),
"image/avif" => Some("avif"),
_ => None,
}
}
fn validate_http_url(value: &str, field: &str) -> Result<()> {
@@ -116,6 +261,54 @@ fn validate_http_url(value: &str, field: &str) -> Result<()> {
Ok(())
}
pub async fn image(
Path(file_name): Path<String>,
) -> std::result::Result<Response<Body>, StatusCode> {
let (hash, extension) = file_name.rsplit_once('.').ok_or(StatusCode::NOT_FOUND)?;
if hash.len() != 64
|| !hash.bytes().all(|byte| byte.is_ascii_hexdigit())
|| !IMAGE_EXTENSIONS.contains(&extension)
{
return Err(StatusCode::NOT_FOUND);
}
let path: PathBuf = config::managed_data_dir()
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
.join("ads")
.join(&file_name);
let bytes = tokio::fs::read(path).await.map_err(|error| {
if error.kind() == std::io::ErrorKind::NotFound {
StatusCode::NOT_FOUND
} else {
StatusCode::INTERNAL_SERVER_ERROR
}
})?;
let mut response = Response::new(Body::from(bytes));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static(content_type_for_extension(extension)),
);
response.headers_mut().insert(
header::CACHE_CONTROL,
HeaderValue::from_static("public, max-age=31536000, immutable"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
Ok(response)
}
fn content_type_for_extension(extension: &str) -> &'static str {
match extension {
"png" => "image/png",
"jpg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
"avif" => "image/avif",
_ => "application/octet-stream",
}
}
pub async fn get(
State(service): State<ControlService>,
headers: HeaderMap,
@@ -146,3 +339,70 @@ pub async fn dismiss(
service.dismiss_ad(&ad_id, &input).await?;
Ok(StatusCode::NO_CONTENT)
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
use axum::{
body::Body,
http::{header, Response},
routing::get,
Router,
};
use super::*;
#[test]
fn recognizes_supported_image_content_types() {
assert_eq!(image_extension("image/png"), Some("png"));
assert_eq!(image_extension("image/jpeg; charset=binary"), Some("jpg"));
assert_eq!(image_extension("IMAGE/WEBP"), Some("webp"));
assert_eq!(image_extension("image/svg+xml"), None);
assert_eq!(image_extension("text/html"), None);
}
#[tokio::test]
async fn downloads_an_ad_image_once_and_reuses_the_cache() {
let requests = Arc::new(AtomicUsize::new(0));
let request_counter = requests.clone();
let app = Router::new().route(
"/ad.png",
get(move || {
request_counter.fetch_add(1, Ordering::SeqCst);
async {
Response::builder()
.header(header::CONTENT_TYPE, "image/png")
.body(Body::from(&b"cached image"[..]))
.unwrap()
}
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
let root = tempfile::tempdir().unwrap();
let client = reqwest::Client::new();
let remote_url = format!("http://{address}/ad.png");
let first = cache_image(&client, root.path(), &remote_url)
.await
.unwrap();
let second = cache_image(&client, root.path(), &remote_url)
.await
.unwrap();
assert_eq!(first, second);
assert!(first.starts_with(ADS_IMAGE_ROUTE));
assert_eq!(requests.load(Ordering::SeqCst), 1);
let file_name = first.rsplit('/').next().unwrap();
assert_eq!(
tokio::fs::read(root.path().join(file_name)).await.unwrap(),
b"cached image"
);
server.abort();
}
}
+5 -8
View File
@@ -112,6 +112,7 @@ fn proxy_error(error: impl std::fmt::Display) -> Response<Body> {
pub fn api_router(service: ControlService) -> Router {
Router::new()
.route("/__byok-api__/api/ads", get(ads::get))
.route("/__byok-api__/api/ads/images/{file_name}", get(ads::image))
.route(
"/__byok-api__/api/ads/{ad_id}/dismissals",
post(ads::dismiss),
@@ -138,14 +139,6 @@ 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/disabled-models",
get(plugins::get_disabled_models).put(plugins::set_disabled_models),
)
.route(
"/__byok-api__/api/plugins/disabled-accounts",
get(plugins::get_disabled_accounts).put(plugins::set_disabled_accounts),
)
.route(
"/__byok-api__/api/plugins/runtime",
get(plugins::runtime_status)
@@ -208,6 +201,10 @@ pub fn api_router(service: ControlService) -> Router {
"/__byok-api__/api/settings/desktop",
get(settings::get_desktop).put(settings::update_desktop),
)
.route(
"/__byok-api__/api/settings/commit",
get(settings::get_commit).put(settings::update_commit),
)
.route(
"/__byok-api__/api/harness/cursor/status",
get(harness::status),
+7 -1
View File
@@ -16,6 +16,7 @@ pub struct OverviewRange {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<String>,
bucket_ms: Option<i64>,
}
pub async fn get(
@@ -24,7 +25,12 @@ pub async fn get(
) -> Result<Json<Overview>> {
Ok(Json(
service
.overview(range.start_ms, range.end_ms, range.model_hashes.as_deref())
.overview(
range.start_ms,
range.end_ms,
range.model_hashes.as_deref(),
range.bucket_ms,
)
.await?,
))
}
-43
View File
@@ -4,7 +4,6 @@ use axum::{
http::StatusCode,
Json,
};
use serde::Deserialize;
use crate::{
plugin::{
@@ -123,45 +122,3 @@ pub async fn cancel_runtime_initialization(
) -> Result<Json<PluginRuntimeStatus>> {
Ok(Json(service.cancel_plugin_runtime_initialization()))
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SetDisabledModelsInput {
pub model_ids: Vec<String>,
}
pub async fn get_disabled_models(
State(service): State<ControlService>,
) -> Result<Json<Vec<String>>> {
Ok(Json(service.disabled_plugin_models().await?))
}
pub async fn set_disabled_models(
State(service): State<ControlService>,
Json(input): Json<SetDisabledModelsInput>,
) -> Result<Json<Vec<String>>> {
service.set_disabled_plugin_models(input.model_ids).await?;
Ok(Json(service.disabled_plugin_models().await?))
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SetDisabledAccountsInput {
pub account_ids: Vec<String>,
}
pub async fn get_disabled_accounts(
State(service): State<ControlService>,
) -> Result<Json<Vec<String>>> {
Ok(Json(service.disabled_plugin_accounts().await?))
}
pub async fn set_disabled_accounts(
State(service): State<ControlService>,
Json(input): Json<SetDisabledAccountsInput>,
) -> Result<Json<Vec<String>>> {
service
.set_disabled_plugin_accounts(input.account_ids)
.await?;
Ok(Json(service.disabled_plugin_accounts().await?))
}
+38 -32
View File
@@ -28,8 +28,8 @@ use crate::{
plugin::{PluginDescriptor, PluginRegistry, PluginRuntime, PluginRuntimeStatus},
provider::{is_valid_response_event, ModelEvent, Provider},
store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store,
TabSettings,
CommitSettings, DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput,
StatisticsStorage, Store, TabSettings,
},
Error, Result,
};
@@ -41,6 +41,8 @@ pub struct ControlService {
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
app_version: String,
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
}
@@ -151,6 +153,8 @@ impl ControlService {
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
app_version: String,
) -> Result<Self> {
Ok(Self {
cursor_harness: CursorHarness::new(store.clone())?,
@@ -158,6 +162,8 @@ impl ControlService {
provider,
plugin_runtime,
plugins,
clients,
app_version,
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
})
}
@@ -251,40 +257,18 @@ impl ControlService {
self.plugin_runtime.cancel_initialization()
}
pub async fn disabled_plugin_models(&self) -> Result<Vec<String>> {
let mut list: Vec<_> = self.store.disabled_plugin_models().await?.into_iter().collect();
list.sort();
Ok(list)
}
pub async fn set_disabled_plugin_models(&self, model_ids: Vec<String>) -> Result<()> {
let set = model_ids.into_iter().collect();
self.store.set_disabled_plugin_models(&set).await
}
pub async fn disabled_plugin_accounts(&self) -> Result<Vec<String>> {
let mut list: Vec<_> = self.store.disabled_plugin_accounts().await?.into_iter().collect();
list.sort();
Ok(list)
}
pub async fn set_disabled_plugin_accounts(&self, account_ids: Vec<String>) -> Result<()> {
let set = account_ids.into_iter().collect();
self.store.set_disabled_plugin_accounts(&set).await
}
pub(super) async fn ads(
&self,
disabled_ad_ids: Option<&str>,
language: &str,
) -> Result<AdRuntime> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let installation_id = self.store.installation_id().await?;
let mut request = client
.get(ADS_ENDPOINT)
.header(DEVICE_ID_HEADER, installation_id)
.header(OS_HEADER, std::env::consts::OS)
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
.header(APP_VERSION_HEADER, &self.app_version)
.header(LANGUAGE_HEADER, language)
.timeout(std::time::Duration::from_secs(60));
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
@@ -299,11 +283,13 @@ impl ControlService {
message.chars().take(200).collect::<String>()
)));
}
response.json::<AdRuntime>().await?.into_menu_slots()
let mut runtime = response.json::<AdRuntime>().await?.into_menu_slots()?;
runtime.cache_images(&client).await;
Ok(runtime)
}
pub(super) async fn dismiss_ad(&self, ad_id: &str, input: &AdDismissalInput) -> Result<()> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let installation_id = self.store.installation_id().await?;
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
Error::Config(format!("advertisement endpoint is invalid: {error}"))
@@ -318,7 +304,7 @@ impl ControlService {
.post(endpoint)
.header(DEVICE_ID_HEADER, installation_id)
.header(OS_HEADER, std::env::consts::OS)
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
.header(APP_VERSION_HEADER, &self.app_version)
.json(input)
.timeout(std::time::Duration::from_secs(5))
.send()
@@ -343,8 +329,11 @@ impl ControlService {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<&str>,
bucket_ms: Option<i64>,
) -> Result<Overview> {
self.store.overview(start_ms, end_ms, model_hashes).await
self.store
.overview(start_ms, end_ms, model_hashes, bucket_ms)
.await
}
pub async fn create_models(&self, models: &[ModelConfigInput]) -> Result<Vec<ModelConfig>> {
@@ -532,7 +521,7 @@ impl ControlService {
}
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let base_url = crate::model::normalize_request_url(&input.base_url)?;
discover_models_from_endpoint(
&client,
@@ -699,7 +688,16 @@ impl ControlService {
}
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
self.store.set_proxy_settings(settings).await
if settings.mode.is_custom() {
let local_proxy_port = match self.cursor_harness.proxy_port().await {
Some(port) => port,
None => self.store.port_settings().await?.proxy_port,
};
crate::network::reject_self_proxy(&settings.address, local_proxy_port)?;
}
let settings = self.store.set_proxy_settings(settings).await?;
self.clients.invalidate().await;
Ok(settings)
}
pub async fn tab_settings(&self) -> Result<TabSettings> {
@@ -717,6 +715,14 @@ impl ControlService {
pub async fn set_desktop_settings(&self, settings: DesktopSettings) -> Result<()> {
self.store.set_desktop_settings(settings).await
}
pub async fn commit_settings(&self) -> Result<CommitSettings> {
self.store.commit_settings().await
}
pub async fn set_commit_settings(&self, settings: CommitSettings) -> Result<CommitSettings> {
self.store.set_commit_settings(settings).await
}
}
fn official_call(trace: CursorRunTraceSummary) -> CallSummary {
+75 -4
View File
@@ -1,11 +1,15 @@
//! Implements settings management endpoints.
use crate::Result;
use axum::{extract::State, Json};
use serde::Deserialize;
use axum::{
extract::State,
http::{header, HeaderMap},
Json,
};
use serde::{Deserialize, Serialize};
use crate::store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage,
StatisticsStorageScope, TabSettings,
CommitPromptLocale, CommitSettings, DesktopSettings, PortSettings, ProxySettings,
ProxySettingsInput, StatisticsStorage, StatisticsStorageScope, TabSettings,
};
use super::{ControlService, ObservabilitySettings};
@@ -87,3 +91,70 @@ pub async fn update_desktop(
service.set_desktop_settings(settings).await?;
get_desktop(State(service)).await
}
/// Settings view for commit message generation. Empty `model_id` means 直连
/// (forward the original Cursor RPC). A non-empty value is a configured
/// built-in or plugin model identifier. Empty `prompt` means "use the built-in default".
#[derive(Serialize)]
pub struct CommitSettingsView {
pub model_id: String,
pub prompt: String,
pub prompt_locale: CommitPromptLocale,
pub default_prompt: &'static str,
}
impl CommitSettingsView {
fn new(settings: CommitSettings, default_locale: CommitPromptLocale) -> Self {
Self {
model_id: settings.model_id,
prompt: settings.prompt,
prompt_locale: settings.prompt_locale,
default_prompt: default_locale.default_prompt(),
}
}
}
pub async fn get_commit(
State(service): State<ControlService>,
headers: HeaderMap,
) -> Result<Json<CommitSettingsView>> {
let settings = service.commit_settings().await?;
Ok(Json(CommitSettingsView::new(
settings,
requested_commit_locale(&headers),
)))
}
pub async fn update_commit(
State(service): State<ControlService>,
Json(settings): Json<CommitSettings>,
) -> Result<Json<CommitSettingsView>> {
let saved = service.set_commit_settings(settings).await?;
let default_locale = saved.prompt_locale;
Ok(Json(CommitSettingsView::new(saved, default_locale)))
}
fn requested_commit_locale(headers: &HeaderMap) -> CommitPromptLocale {
match headers
.get(header::ACCEPT_LANGUAGE)
.and_then(|value| value.to_str().ok())
{
Some(value) if value.eq_ignore_ascii_case("zh-CN") => CommitPromptLocale::ZhCn,
_ => CommitPromptLocale::EnUs,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn commit_default_prompt_locale_comes_from_interface_language() {
let mut headers = HeaderMap::new();
headers.insert(header::ACCEPT_LANGUAGE, "zh-CN".parse().unwrap());
assert_eq!(requested_commit_locale(&headers), CommitPromptLocale::ZhCn);
headers.insert(header::ACCEPT_LANGUAGE, "en-US".parse().unwrap());
assert_eq!(requested_commit_locale(&headers), CommitPromptLocale::EnUs);
}
}
+11 -13
View File
@@ -278,19 +278,17 @@ impl CheckpointBuilder {
),
});
if let Some(trace) = handle.trace() {
trace
.artifact(
"checkpoint",
"byok_server",
&checkpoint.encode_to_vec(),
serde_json::json!({
"root_message_count": checkpoint.root_prompt_messages_json.len(),
"turn_count": checkpoint.turns.len(),
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
"emit_status": if result.is_ok() { "sent" } else { "error" },
}),
)
.await;
trace.artifact(
"checkpoint",
"byok_server",
&checkpoint.encode_to_vec(),
serde_json::json!({
"root_message_count": checkpoint.root_prompt_messages_json.len(),
"turn_count": checkpoint.turns.len(),
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
"emit_status": if result.is_ok() { "sent" } else { "error" },
}),
);
}
result
}
+55 -3
View File
@@ -4,7 +4,10 @@ use std::collections::HashMap;
use prost::Message;
use crate::{
cursor::{prompting::fold_derived_state, protocol::proto::agent::v1 as pb},
cursor::{
prompting::{fold_derived_state, fold_derived_state_from, DerivedState},
protocol::proto::agent::v1 as pb,
},
model::{CanonicalMessage, MessageContent},
store::BlobId,
Error, Result,
@@ -17,7 +20,18 @@ impl CheckpointBuilder {
&self,
messages: &[CanonicalMessage],
) -> Result<(Vec<BlobId>, Option<BlobId>)> {
let state = fold_derived_state(messages);
let changes = fold_derived_state(messages);
let state = if changes.todos.is_some() && !self.base.todos.is_empty() {
fold_derived_state_from(
messages,
DerivedState {
todos: Some(self.base_todo_state().await?),
plan: None,
},
)
} else {
changes
};
let todo_values = state
.todos
.as_ref()
@@ -28,7 +42,15 @@ impl CheckpointBuilder {
.ok_or_else(|| Error::Protocol("TodoWrite state is missing todos[]".into()))
})
.transpose()?;
let mut todo_ids = Vec::new();
let mut todo_ids = if todo_values.is_none() {
self.base
.todos
.iter()
.map(|raw| BlobId::from_bytes(raw))
.collect::<Result<Vec<_>>>()?
} else {
Vec::new()
};
for (index, todo) in todo_values.into_iter().flatten().enumerate() {
let status = match todo
.get("status")
@@ -97,6 +119,36 @@ impl CheckpointBuilder {
};
Ok((todo_ids, plan_id))
}
async fn base_todo_state(&self) -> Result<serde_json::Value> {
let mut todos = Vec::with_capacity(self.base.todos.len());
for raw_id in &self.base.todos {
let id = BlobId::from_bytes(raw_id)?;
let data = self.sync.get(&id).await?.ok_or_else(|| {
Error::Protocol(format!("Cursor Todo Blob is missing: {}", id.to_base64()))
})?;
let todo = pb::TodoItem::decode(data.as_slice())?;
let status = match pb::TodoStatus::try_from(todo.status) {
Ok(pb::TodoStatus::InProgress) => "in_progress",
Ok(pb::TodoStatus::Completed) => "completed",
Ok(pb::TodoStatus::Cancelled) => "cancelled",
Ok(pb::TodoStatus::Pending) => "pending",
Ok(pb::TodoStatus::Unspecified) | Err(_) => {
return Err(Error::Protocol(format!(
"unknown Cursor Todo status: {}",
todo.status
)))
}
};
todos.push(serde_json::json!({
"id": todo.id,
"content": todo.content,
"status": status,
"dependencies": todo.dependencies,
}));
}
Ok(serde_json::json!({"merge": false, "todos": todos}))
}
}
pub(super) fn update_current_step_state(
+2 -14
View File
@@ -12,7 +12,7 @@ use crate::{
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
tools::runtime::McpRoute,
},
model::ToolDefinition,
model::{normalize_tool_name, ToolDefinition},
store::BlobId,
Error, Result,
};
@@ -493,7 +493,7 @@ pub fn dynamic_mcp(
})?),
};
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
let name = model_tool_name(&wire.name);
let name = normalize_tool_name(&wire.name);
let definition = ToolDefinition {
name: name.clone(),
description: wire.description.clone(),
@@ -553,18 +553,6 @@ fn invalid_mcp_parameters(tool_name: &str) -> Error {
))
}
fn model_tool_name(name: &str) -> String {
name.chars()
.map(|character| {
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
character
} else {
'_'
}
})
.collect()
}
fn prost_value(value: &prost_types::Value) -> Value {
use prost_types::value::Kind;
match value.kind.as_ref() {
+24 -5
View File
@@ -119,11 +119,7 @@ fn from_requested(
))
})?);
}
other => {
return Err(Error::Protocol(format!(
"unsupported Cursor model parameter: {other}"
)))
}
_ => {}
}
}
Ok(spec)
@@ -139,3 +135,26 @@ fn parse_bool(parameter: &pb::requested_model::ModelParameterValue) -> Result<bo
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ignores_unknown_cursor_model_parameters() {
let requested = pb::RequestedModel {
model_id: "test-model".into(),
parameters: vec![pb::requested_model::ModelParameterValue {
id: "optimize_for".into(),
value: "quality".into(),
}],
..Default::default()
};
let model = from_requested(&requested, None).expect("unknown parameter should be ignored");
assert_eq!(model.model_id, "test-model");
assert_eq!(model.latency, ModelLatency::Standard);
assert!(!model.reasoning.enabled);
}
}
+1 -3
View File
@@ -117,9 +117,7 @@ pub(crate) async fn prepare(
"selected_source": "root_prompt_messages_json",
});
let encoded = serde_json::to_vec(&summary)?;
trace
.artifact("history_projection", "byok_server", &encoded, summary)
.await;
trace.artifact("history_projection", "byok_server", &encoded, summary);
}
let mut request_context = context::hydrate(request, context_sync).await?;
if let Some(rules_dir) = local_rules_dir {
+18 -2
View File
@@ -1,6 +1,19 @@
//! Defines commands accepted by a Conversation runtime.
use crate::cursor::protocol::proto::agent::v1 as pb;
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
#[derive(Debug)]
pub enum RunFinish {
TurnCompleted,
Transport(TransportFinish),
}
#[derive(Debug)]
pub enum TransportFinish {
Success,
Failed(Error),
Cancelled,
}
#[derive(Debug)]
pub enum TransportCommand {
@@ -8,6 +21,9 @@ pub enum TransportCommand {
seqno: i64,
message: Box<pb::AgentClientMessage>,
},
RunFinished {
generation: u64,
finish: RunFinish,
},
Disconnect,
Close,
}
+31 -12
View File
@@ -34,7 +34,7 @@ use crate::{
Error, Result,
};
use super::{CompiledMessages, ConversationRegistry, MessageDelivery};
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
use crate::cursor::transport::TransportHandle;
pub struct ConversationOutput {
@@ -110,7 +110,7 @@ impl ConversationOutput {
}
}
pub async fn run(mut self) -> Result<()> {
pub async fn run(mut self) -> Result<RunFinish> {
let result = self.run_inner().await;
if let Err(error) = &result {
if !self.superseded.is_cancelled() {
@@ -143,7 +143,7 @@ impl ConversationOutput {
result
}
async fn run_inner(&mut self) -> Result<()> {
async fn run_inner(&mut self) -> Result<RunFinish> {
if self.context.compacting {
self.handle.emit(&events::summary_started())?;
}
@@ -158,6 +158,7 @@ impl ConversationOutput {
let mut streams = BTreeMap::<usize, ToolCallStream>::new();
let mut completions = HashMap::<String, ToolCompletion>::new();
let mut completed = HashSet::<String>::new();
let mut completed_round = None::<ToolRoundId>;
let mut response_text = String::new();
let mut response_thinking = String::new();
let mut active_round = None::<ToolRoundId>;
@@ -175,7 +176,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
let input = if let Ok(action) = self.runtime_actions.try_recv() {
Input::RuntimeAction(Some(Box::new(action)))
@@ -187,7 +188,7 @@ impl ConversationOutput {
_ = self.superseded.cancelled() => {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
event = self.core.events.recv() => Input::Event(event),
@@ -399,6 +400,14 @@ impl ConversationOutput {
.unwrap_or_else(|_| serde_json::json!({}))
};
}
RunEvent::UsageSnapshot(usage) => {
if !self.context.compacting {
if let Some(output_tokens) = usage.output_tokens {
self.handle.emit(&events::token_delta(output_tokens))?;
}
context_tokens = usage.context_input_tokens;
}
}
RunEvent::Usage(usage) => {
if !self.context.compacting {
if let Some(output_tokens) = usage.output_tokens {
@@ -417,6 +426,16 @@ impl ConversationOutput {
round_id,
calls: round_calls,
} => {
// `completed` exists so that replaying ExecuteToolRound for a round
// does not dispatch a call this round already committed. Tool call ids
// are only unique *within* a round -- the schema says as much with
// `UNIQUE (round_id, call_id)` -- so an id retained from an earlier
// round would make start_batch skip a fresh call, and tool_round wait
// forever for a result nothing will ever produce.
if completed_round.as_ref() != Some(&round_id) {
completed.clear();
completed_round = Some(round_id.clone());
}
active_round = Some(round_id.clone());
active_tool_calls = round_calls
.iter()
@@ -711,7 +730,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
return match outcome {
RunOutcome::Completed => {
@@ -727,8 +746,7 @@ impl ConversationOutput {
for _ in 0..3 {
self.checkpoint.publish(&self.handle, &checkpoint).await?;
}
finish_success(&self.handle);
return Ok(());
return Ok(RunFinish::TurnCompleted);
}
let checkpoints = final_checkpoint.take().ok_or_else(|| {
Error::Protocol("Completed without final state".into())
@@ -744,18 +762,19 @@ impl ConversationOutput {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
})?;
finish_success(&self.handle);
Ok(())
Ok(RunFinish::TurnCompleted)
}
RunOutcome::Cancelled => {
worker.abort();
self.abort_execs().await;
finish_cancelled(&self.handle)
Ok(RunFinish::Transport(TransportFinish::Cancelled))
}
RunOutcome::Failed(failure) => {
worker.abort();
self.abort_execs().await;
finish_failed(&self.handle, &cursor_error(failure))
Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error(
failure,
))))
}
};
}
+287 -94
View File
@@ -18,18 +18,22 @@ use crate::{
},
transport::{OrderedInbox, TransportHandle},
},
run::{CommandResult, RunEngine, RunHandle},
run::{CommandResult, RunEngine, RunHandle, RunPhase},
};
use super::{
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
ConversationRegistry, MessageDelivery, TransportCommand,
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
};
pub struct ConversationRuntime;
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
#[derive(Clone)]
struct RunGeneration {
id: u64,
request: pb::AgentRunRequest,
superseded: CancellationToken,
finished: CancellationToken,
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
@@ -41,6 +45,14 @@ struct RunGeneration {
struct FinishGeneration(CancellationToken);
struct TransportActorGuard(TransportHandle);
impl Drop for TransportActorGuard {
fn drop(&mut self) {
self.0.close_transport();
}
}
impl Drop for FinishGeneration {
fn drop(&mut self) {
self.0.cancel();
@@ -54,6 +66,7 @@ impl ConversationRuntime {
mut receiver: mpsc::Receiver<TransportCommand>,
) {
tokio::spawn(async move {
let _actor_guard = TransportActorGuard(handle.clone());
let dependencies = registry.dependencies().clone();
let blob_sync = BlobSynchronizer::new(
handle.request_id().into(),
@@ -65,19 +78,70 @@ impl ConversationRuntime {
let context_sync =
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
let mut current = None::<RunGeneration>;
let mut next_generation = 1_u64;
let mut pending_finish = None::<(u64, TransportFinish)>;
let mut draining = false;
let mut waiting_for_action = false;
loop {
let command = match receiver.recv().await {
Some(command) => command,
None => {
handle.mark_disconnected();
if let Some(generation) = current.as_ref() {
generation.superseded.cancel();
if let Some(run) = generation.run.lock().clone() {
run.cancel();
let command = if draining {
if !handle.admissions_drained() {
tokio::select! {
command = receiver.recv() => match command {
Some(command) => command,
None => {
finish_pending(&handle, &current, pending_finish.take());
break;
}
},
_ = handle.wait_admissions_drained() => continue,
}
} else {
handle.mark_draining();
match receiver.try_recv() {
Ok(command) => command,
Err(mpsc::error::TryRecvError::Empty)
| Err(mpsc::error::TryRecvError::Disconnected) => {
finish_pending(&handle, &current, pending_finish.take());
break;
}
}
super::finish_cancelled(&handle).ok();
break;
}
} else if waiting_for_action {
tokio::select! {
command = receiver.recv() => match command {
Some(command) => command,
None => {
handle.mark_disconnected();
super::finish_success(&handle);
break;
}
},
_ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => {
let Some(generation) = current.as_ref() else {
super::finish_success(&handle);
break;
};
handle.begin_close();
pending_finish = Some((generation.id, TransportFinish::Success));
draining = true;
waiting_for_action = false;
continue;
}
}
} else {
match receiver.recv().await {
Some(command) => command,
None => {
handle.mark_disconnected();
if let Some(generation) = current.as_ref() {
generation.superseded.cancel();
if let Some(run) = generation.run.lock().clone() {
run.cancel();
}
}
super::finish_cancelled(&handle).ok();
break;
}
}
};
match command {
@@ -92,11 +156,41 @@ impl ConversationRuntime {
let _ = handle.emit(&codec::abort(id));
}
}
super::finish_cancelled(&handle).ok();
let turn_completed = waiting_for_action
|| current.as_ref().is_some_and(|generation| {
generation
.run
.lock()
.as_ref()
.is_none_or(|run| run.phase() != RunPhase::Running)
});
if turn_completed {
super::finish_success(&handle);
} else {
super::finish_cancelled(&handle).ok();
}
break;
}
TransportCommand::Close => {
break;
TransportCommand::RunFinished { generation, finish } => {
if current
.as_ref()
.is_none_or(|current| current.id != generation)
{
continue;
}
match finish {
RunFinish::TurnCompleted => {
pending_finish = None;
draining = false;
waiting_for_action = true;
}
RunFinish::Transport(finish) => {
waiting_for_action = false;
handle.begin_close();
pending_finish = Some((generation, finish));
draining = true;
}
}
}
TransportCommand::Append { seqno, message } => {
for (_seqno, message) in inbox.push(seqno, *message) {
@@ -105,6 +199,12 @@ impl ConversationRuntime {
Some(pb::agent_client_message::Message::RunRequest(
request,
)) => {
waiting_for_action = false;
if draining {
handle.reopen();
draining = false;
pending_finish = None;
}
if let Some(conversation_id) =
request.conversation_id.as_deref()
{
@@ -117,60 +217,21 @@ impl ConversationRuntime {
"invalid Cursor conversation id"
);
let _ = super::finish_failed(&handle, &error);
let _ =
handle.command(TransportCommand::Close).await;
return;
}
}
let previous_finished =
if let Some(previous) = current.take() {
previous.superseded.cancel();
if let Some(run) = previous.run.lock().clone() {
run.cancel();
}
for id in previous
.tool_runtime
.interrupt_for_run_replacement()
.await
{
let _ = handle.emit(&codec::abort(id));
}
Some(previous.finished.clone())
} else {
None
};
let (results, result_receiver) = tool_result_channel();
let (runtime_actions, runtime_action_receiver) =
mpsc::unbounded_channel::<compile::RuntimeAction>();
let tool_runtime = tool_runtime_factory.next_run();
let tools = ToolDispatcher::with_results(
tool_runtime.clone(),
results.clone(),
dependencies.store.clone(),
dependencies.web_cache.clone(),
);
let generation = RunGeneration {
superseded: CancellationToken::new(),
finished: CancellationToken::new(),
run: Arc::new(parking_lot::Mutex::new(None)),
results,
runtime_actions,
tool_runtime,
tools,
};
current = Some(generation.clone());
spawn_run_request(
registry.clone(),
handle.clone(),
start_generation(
&registry,
&handle,
&dependencies,
&blob_sync,
&context_sync,
&tool_runtime_factory,
&mut current,
&mut next_generation,
request,
dependencies.clone(),
blob_sync.clone(),
context_sync.clone(),
generation,
previous_finished,
result_receiver,
runtime_action_receiver,
);
)
.await;
}
Some(pb::agent_client_message::Message::ExecClientMessage(
message,
@@ -331,27 +392,48 @@ impl ConversationRuntime {
// return an explicit Protocol Error rather than falling through silently.
Some(
pb::agent_client_message::Message::ConversationAction(
action,
conversation_action,
),
) => match action.action {
) => match conversation_action.action.clone() {
Some(
pb::conversation_action::Action::UserMessageAction(
action,
),
) => {
let Some(generation) = current.as_ref() else {
let delivered_to_active_run =
current.as_ref().is_some_and(|generation| {
generation.run.lock().as_ref().is_some_and(
|run| run.phase() == RunPhase::Running,
) && generation
.runtime_actions
.send(compile::RuntimeAction::UserMessage(
action.clone(),
))
.is_ok()
});
if delivered_to_active_run {
continue;
}
let Some(previous) = current.as_ref() else {
continue;
};
if generation
.runtime_actions
.send(compile::RuntimeAction::UserMessage(action))
.is_err()
{
generation.results.send_error(crate::Error::Protocol(
"UserMessageAction arrived without an active Run"
.into(),
));
}
let mut request = previous.request.clone();
request.action = Some(conversation_action);
request.conversation_state = None;
request.pre_fetched_blobs.clear();
waiting_for_action = false;
start_generation(
&registry,
&handle,
&dependencies,
&blob_sync,
&context_sync,
&tool_runtime_factory,
&mut current,
&mut next_generation,
request,
)
.await;
}
Some(pb::conversation_action::Action::CancelAction(_)) => {
if let Some(generation) = current.as_ref() {
@@ -428,6 +510,95 @@ impl ConversationRuntime {
}
}
#[allow(clippy::too_many_arguments)]
async fn start_generation(
registry: &ConversationRegistry,
handle: &TransportHandle,
dependencies: &ConversationDependencies,
blob_sync: &BlobSynchronizer,
context_sync: &RequestContextSynchronizer,
tool_runtime_factory: &CursorToolRuntime,
current: &mut Option<RunGeneration>,
next_generation: &mut u64,
request: pb::AgentRunRequest,
) {
let previous_finished = if let Some(previous) = current.take() {
previous.superseded.cancel();
if let Some(run) = previous.run.lock().clone() {
run.cancel();
}
for id in previous.tool_runtime.interrupt_for_run_replacement().await {
let _ = handle.emit(&codec::abort(id));
}
Some(previous.finished.clone())
} else {
None
};
let (results, result_receiver) = tool_result_channel();
let (runtime_actions, runtime_action_receiver) =
mpsc::unbounded_channel::<compile::RuntimeAction>();
let tool_runtime = tool_runtime_factory.next_run();
let tools = ToolDispatcher::with_results(
tool_runtime.clone(),
results.clone(),
dependencies.store.clone(),
dependencies.web_cache.clone(),
);
let generation = RunGeneration {
id: *next_generation,
request: request.clone(),
superseded: CancellationToken::new(),
finished: CancellationToken::new(),
run: Arc::new(parking_lot::Mutex::new(None)),
results,
runtime_actions,
tool_runtime,
tools,
};
*next_generation = next_generation.saturating_add(1);
*current = Some(generation.clone());
spawn_run_request(
registry.clone(),
handle.clone(),
request,
dependencies.clone(),
blob_sync.clone(),
context_sync.clone(),
generation,
previous_finished,
result_receiver,
runtime_action_receiver,
);
}
fn finish_pending(
handle: &TransportHandle,
current: &Option<RunGeneration>,
pending: Option<(u64, TransportFinish)>,
) {
let Some((generation, finish)) = pending else {
return;
};
if current
.as_ref()
.is_some_and(|current| current.id == generation)
{
finish_transport(handle, finish);
}
}
fn finish_transport(handle: &TransportHandle, finish: TransportFinish) {
match finish {
TransportFinish::Success => super::finish_success(handle),
TransportFinish::Failed(error) => {
let _ = super::finish_failed(handle, &error);
}
TransportFinish::Cancelled => {
let _ = super::finish_cancelled(handle);
}
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_run_request(
registry: ConversationRegistry,
@@ -487,8 +658,12 @@ fn spawn_run_request(
%error,
"failed to prepare Cursor Run"
);
let _ = super::finish_failed(&handle, &error);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Failed(error)),
})
.await;
return;
}
};
@@ -520,8 +695,12 @@ fn spawn_run_request(
{
CommandResult::Applied | CommandResult::Duplicate => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Success),
})
.await;
}
return;
}
@@ -545,8 +724,12 @@ fn spawn_run_request(
}
CommandResult::StaleTarget => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Success),
})
.await;
}
return;
}
@@ -605,16 +788,21 @@ fn spawn_run_request(
tool_runtime: generation.tool_runtime.clone(),
},
);
if let Err(error) = output.run().await {
if !generation.superseded.is_cancelled() {
tracing::error!(
request_id = handle.request_id(),
%error,
"Cursor session failed"
);
let _ = super::finish_failed(&handle, &error);
let finish = match output.run().await {
Ok(finish) => finish,
Err(error) => {
if generation.superseded.is_cancelled() {
RunFinish::Transport(TransportFinish::Cancelled)
} else {
tracing::error!(
request_id = handle.request_id(),
%error,
"Cursor session failed"
);
RunFinish::Transport(TransportFinish::Failed(error))
}
}
}
};
let _ = core_run.await;
registry.release(&conversation_id, &run_id).await;
if generation
@@ -626,7 +814,12 @@ fn spawn_run_request(
*generation.run.lock() = None;
}
if !generation.superseded.is_cancelled() {
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish,
})
.await;
}
});
}
+86 -1
View File
@@ -11,7 +11,13 @@ pub struct DerivedState {
}
pub fn fold_derived_state(messages: &[CanonicalMessage]) -> DerivedState {
let mut state = DerivedState::default();
fold_derived_state_from(messages, DerivedState::default())
}
pub fn fold_derived_state_from(
messages: &[CanonicalMessage],
mut state: DerivedState,
) -> DerivedState {
let mut calls = std::collections::HashMap::<String, (String, Value)>::new();
for message in messages {
match &message.content {
@@ -81,3 +87,82 @@ fn normalize(value: &str) -> String {
.flat_map(char::to_lowercase)
.collect()
}
#[cfg(test)]
mod tests {
use serde_json::{json, Value};
use super::{fold_derived_state_from, DerivedState};
use crate::model::{
CanonicalMessage, MessageContent, Origin, Role, ToolCallContent, ToolResultContent,
};
#[test]
fn merge_patch_inherits_content_from_checkpoint_todo_state() {
let messages = todo_write_messages(json!({
"merge": true,
"todos": [{"id": "tests", "status": "completed"}],
}));
let initial = DerivedState {
todos: Some(json!({
"merge": false,
"todos": [{
"id": "tests",
"content": "Run focused tests",
"status": "in_progress",
}],
})),
plan: None,
};
let state = fold_derived_state_from(&messages, initial);
assert_eq!(
state.todos,
Some(json!({
"merge": false,
"todos": [{
"id": "tests",
"content": "Run focused tests",
"status": "completed",
}],
}))
);
}
fn todo_write_messages(arguments: Value) -> Vec<CanonicalMessage> {
vec![
CanonicalMessage {
message_id: "assistant".into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: String::new(),
thinking: String::new(),
tool_round_id: None,
replay_state: None,
tool_calls: vec![ToolCallContent {
index: 0,
call_id: "todo-call".into(),
name: "TodoWrite".into(),
arguments,
}],
},
runtime_event_id: None,
},
CanonicalMessage {
message_id: "result".into(),
role: Role::Tool,
origin: Origin::Tool,
content: MessageContent::ToolResult(ToolResultContent {
call_id: "todo-call".into(),
name: "TodoWrite".into(),
content: "{}".into(),
is_error: false,
image: None,
provider_parts: Vec::new(),
}),
runtime_event_id: None,
},
]
}
}
+35
View File
@@ -26,9 +26,44 @@ pub mod aiserver {
pub data_binary: Vec<u8>,
}
/// Commit message generation request. Only the fields the local
/// generator consumes are decoded; credentials and heavyweight context
/// fields are intentionally left to prost's unknown-field skipping.
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct WriteGitCommitMessageRequest {
#[prost(string, repeated, tag = "1")]
pub diffs: Vec<String>,
#[prost(string, repeated, tag = "2")]
pub previous_commit_messages: Vec<String>,
#[prost(message, optional, tag = "3")]
pub explicit_context: Option<ExplicitContext>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct ExplicitContext {
#[prost(string, tag = "1")]
pub context: String,
#[prost(string, optional, tag = "2")]
pub repo_context: Option<String>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct WriteGitCommitMessageResponse {
#[prost(string, tag = "1")]
pub commit_message: String,
}
#[derive(Clone, Copy, PartialEq, ::prost::Message)]
pub struct BidiAppendResponse {}
/// `NetworkService/IsConnected` reply. The Cursor extension probes this
/// ~10s after any slow request starts; a non-OK result is treated as
/// "network disconnected" and aborts in-flight work (e.g. commit message
/// generation) even while the model is still streaming, so it always
/// answers as connected.
#[derive(Clone, Copy, PartialEq, ::prost::Message)]
pub struct IsConnectedResponse {}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct CustomErrorDetails {
#[prost(string, tag = "1")]
+275 -80
View File
@@ -1,15 +1,17 @@
//! Implements Cursor account information services.
use axum::{
body::{Body, Bytes},
body::{to_bytes, Body},
extract::Extension,
http::{header, Request, Response},
http::{header, HeaderValue, Request, Response},
};
use prost::Message;
use serde_json::{Map, Value};
use serde_json::Value;
use crate::{api::cursor::proxy, Result};
use crate::{api::cursor::proxy, local_app, Result};
const LOCAL_AUTH_ID: &str = "local_ultra";
use super::entitlement::FreeEntitlementCache;
const LOCAL_AUTH_ID: &str = "cursor-local-user";
const LOCAL_EMAIL: &str = "cursor@ai.com";
const LOCAL_ULTRA_PLAN_INCLUDED_CENTS: i32 = 20_000;
@@ -21,6 +23,20 @@ struct GetEmailResponse {
sign_up_type: i32,
}
#[derive(Clone, PartialEq, Message)]
struct GetUserMetaResponse {
#[prost(string, tag = "1")]
email: String,
#[prost(int32, tag = "2")]
sign_up_type: i32,
#[prost(int64, tag = "3")]
user_id: i64,
#[prost(string, optional, tag = "4")]
workos_id: Option<String>,
#[prost(string, optional, tag = "5")]
profile_picture_url: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct GetMeResponse {
#[prost(string, tag = "1")]
@@ -41,6 +57,8 @@ struct GetMeResponse {
email_domain_type: Option<String>,
#[prost(string, optional, tag = "12")]
country: Option<String>,
#[prost(string, optional, tag = "13")]
profile_picture_url: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
@@ -134,7 +152,7 @@ pub async fn get_email(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetEmailResponse {
email: LOCAL_EMAIL.into(),
sign_up_type: 3,
@@ -143,11 +161,27 @@ pub async fn get_email(
.await
}
pub async fn get_user_meta(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_forward(upstream, request, || {
proto(GetUserMetaResponse {
email: LOCAL_EMAIL.into(),
sign_up_type: 3,
user_id: 1,
workos_id: None,
profile_picture_url: None,
})
})
.await
}
pub async fn get_me(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetMeResponse {
auth_id: LOCAL_AUTH_ID.into(),
user_id: 1,
@@ -158,6 +192,7 @@ pub async fn get_me(
is_enterprise_user: Some(false),
email_domain_type: Some("personal".into()),
country: Some("US".into()),
profile_picture_url: None,
})
})
.await
@@ -167,14 +202,14 @@ pub async fn get_teams(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || proto(Empty {})).await
local_or_forward(upstream, request, || proto(Empty {})).await
}
pub async fn get_user_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetUserProfileResponse {
public_visibility_allowed: Some(true),
max_visibility: Some("PUBLIC".into()),
@@ -183,85 +218,170 @@ pub async fn get_user_profile(
.await
}
pub async fn current_period_usage() -> Result<Response<Body>> {
let now = chrono::Utc::now();
proto(GetCurrentPeriodUsageResponse {
billing_cycle_start: (now - chrono::Duration::days(30)).timestamp_millis(),
billing_cycle_end: (now + chrono::Duration::days(10 * 365)).timestamp_millis(),
plan_usage: Some(PlanUsage {
total_spend: 0,
included_spend: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
limit: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining_bonus: Some(false),
bonus_tooltip: Some("Ultra local account mock is active.".into()),
auto_spend: Some(0),
api_spend: Some(0),
auto_percent_used: Some(0.0),
api_percent_used: Some(0.0),
total_percent_used: Some(0.0),
}),
spend_limit_usage: Some(SpendLimitUsage {
limit_type: "user".into(),
}),
display_threshold: Some(99_999_999),
enabled: true,
display_message: "Ultra plan active".into(),
auto_model_selected_display_message: Some("Ultra plan active".into()),
named_model_selected_display_message: Some("Ultra plan active".into()),
pub async fn current_period_usage(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_confirmed_free_or_forward(upstream, &free_entitlements, request, || {
let now = chrono::Utc::now();
proto(GetCurrentPeriodUsageResponse {
billing_cycle_start: (now - chrono::Duration::days(30)).timestamp_millis(),
billing_cycle_end: (now + chrono::Duration::days(10 * 365)).timestamp_millis(),
plan_usage: Some(PlanUsage {
total_spend: 0,
included_spend: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
limit: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining_bonus: Some(false),
bonus_tooltip: Some("Ultra local account mock is active.".into()),
auto_spend: Some(0),
api_spend: Some(0),
auto_percent_used: Some(0.0),
api_percent_used: Some(0.0),
total_percent_used: Some(0.0),
}),
spend_limit_usage: Some(SpendLimitUsage {
limit_type: "user".into(),
}),
display_threshold: Some(99_999_999),
enabled: true,
display_message: "Ultra plan active".into(),
auto_model_selected_display_message: Some("Ultra plan active".into()),
named_model_selected_display_message: Some("Ultra plan active".into()),
})
})
.await
}
pub async fn usage_limit_status() -> Result<Response<Body>> {
proto(GetUsageLimitStatusAndActiveGrantsResponse {
usage_limit_policy_status: Some(UsageLimitPolicyStatus {
is_in_slow_pool: false,
features: Default::default(),
can_configure_spend_limit: true,
has_pending_request: false,
allowed_model_ids: Vec::new(),
allowed_model_tags: Vec::new(),
}),
pub async fn usage_limit_status(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_confirmed_free_or_forward(upstream, &free_entitlements, request, || {
proto(GetUsageLimitStatusAndActiveGrantsResponse {
usage_limit_policy_status: Some(UsageLimitPolicyStatus {
is_in_slow_pool: false,
features: Default::default(),
can_configure_spend_limit: true,
has_pending_request: false,
allowed_model_ids: Vec::new(),
allowed_model_tags: Vec::new(),
}),
})
})
.await
}
pub async fn stripe_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => {
let mut profile = serde_json::from_slice::<Map<String, Value>>(&response.body)?;
ultra(&mut profile);
Ok(response.with_body(Bytes::from(serde_json::to_vec(&profile)?)))
}
Ok(response) => {
tracing::warn!(status = %response.status, "Cursor account upstream rejected profile; using local Ultra identity");
json(ultra_profile())
}
Err(error) => {
tracing::warn!(%error, "Cursor account upstream unavailable; using local Ultra identity");
json(ultra_profile())
}
if local_app::request_uses_local_cursor_token(request.headers()) {
return local_stripe_profile(request.headers().get(header::ORIGIN).cloned());
}
let request_headers = request.headers().clone();
let origin = request.headers().get(header::ORIGIN).cloned();
let upstream_response = match proxy::forward_buffered(&upstream, request).await {
Ok(response) => response,
Err(error) if free_entitlements.is_confirmed_free(&request_headers) => {
tracing::warn!(%error, "using cached Free entitlement after Stripe upstream failure");
return local_stripe_profile(origin);
}
Err(error) => return Err(error),
};
if !upstream_response.status.is_success() {
if should_fallback_to_cached_free(
upstream_response.status,
&free_entitlements,
&request_headers,
) {
tracing::warn!(
status = %upstream_response.status,
"using cached Free entitlement after Stripe upstream failure"
);
return local_stripe_profile(origin);
}
return Ok(upstream_response.into_response());
}
let Some(membership_type) = membership_type(&upstream_response.body) else {
return Ok(upstream_response.into_response());
};
let observed = free_entitlements.observe_membership(&request_headers, &membership_type);
if observed && membership_type.eq_ignore_ascii_case("free") {
return local_stripe_profile(origin);
}
Ok(upstream_response.into_response())
}
async fn forward_or(
fn should_fallback_to_cached_free(
status: axum::http::StatusCode,
free_entitlements: &FreeEntitlementCache,
headers: &axum::http::HeaderMap,
) -> bool {
status.is_server_error() && free_entitlements.is_confirmed_free(headers)
}
fn membership_type(body: &[u8]) -> Option<String> {
let profile: Value = serde_json::from_slice(body).ok()?;
let membership_type = profile.get("membershipType")?.as_str()?.trim();
(!membership_type.is_empty()).then(|| membership_type.to_owned())
}
fn local_stripe_profile(origin: Option<HeaderValue>) -> Result<Response<Body>> {
let mut response = json(ultra_profile())?;
if let Some(origin) = origin {
response
.headers_mut()
.insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, origin);
response.headers_mut().insert(
header::ACCESS_CONTROL_ALLOW_CREDENTIALS,
HeaderValue::from_static("true"),
);
response
.headers_mut()
.insert(header::VARY, HeaderValue::from_static("Origin"));
}
Ok(response)
}
async fn local_or_confirmed_free_or_forward(
upstream: proxy::CursorProxy,
free_entitlements: &FreeEntitlementCache,
request: Request<Body>,
local: impl FnOnce() -> Result<Response<Body>>,
) -> Result<Response<Body>> {
if local_app::request_uses_local_cursor_token(request.headers())
|| free_entitlements.is_confirmed_free(request.headers())
{
consume_body(request).await?;
return local();
}
proxy::forward(Extension(upstream), request).await
}
async fn local_or_forward(
upstream: proxy::CursorProxy,
request: Request<Body>,
fallback: impl FnOnce() -> Result<Response<Body>>,
local: impl FnOnce() -> Result<Response<Body>>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => Ok(response.into_response()),
Ok(response) => {
tracing::debug!(status = %response.status, "Cursor identity upstream rejected request; using local identity");
fallback()
}
Err(error) => {
tracing::warn!(%error, "Cursor identity upstream unavailable; using local identity");
fallback()
}
if local_app::request_uses_local_cursor_token(request.headers()) {
consume_body(request).await?;
return local();
}
proxy::forward(Extension(upstream), request).await
}
async fn consume_body(request: Request<Body>) -> Result<()> {
to_bytes(request.into_body(), usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok(())
}
fn proto(message: impl Message) -> Result<Response<Body>> {
@@ -289,15 +409,6 @@ fn response(content_type: &'static str, body: Vec<u8>) -> Result<Response<Body>>
Ok(response)
}
fn ultra(profile: &mut Map<String, Value>) {
profile.insert("membershipType".into(), Value::String("ultra".into()));
profile.insert(
"individualMembershipType".into(),
Value::String("ultra".into()),
);
profile.insert("subscriptionStatus".into(), Value::String("active".into()));
}
fn ultra_profile() -> Value {
serde_json::json!({
"membershipType": "ultra",
@@ -310,3 +421,87 @@ fn ultra_profile() -> Value {
"isTeamMember": false
})
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use axum::body::Bytes;
use futures_util::stream;
use super::*;
#[tokio::test]
async fn local_account_response_consumes_request_body_before_replying() {
let polled = Arc::new(AtomicBool::new(false));
let observed = polled.clone();
let body = Body::from_stream(stream::once(async move {
observed.store(true, Ordering::SeqCst);
Ok::<_, std::convert::Infallible>(Bytes::from_static(b"request"))
}));
consume_body(Request::new(body)).await.unwrap();
assert!(polled.load(Ordering::SeqCst));
}
#[test]
fn cached_free_fallback_accepts_server_failures_but_not_auth_failures() {
let cache = FreeEntitlementCache::default();
let mut headers = axum::http::HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
HeaderValue::from_static("Bearer official-free-token"),
);
assert!(cache.observe_membership(&headers, "free"));
assert!(should_fallback_to_cached_free(
axum::http::StatusCode::BAD_GATEWAY,
&cache,
&headers
));
assert!(!should_fallback_to_cached_free(
axum::http::StatusCode::UNAUTHORIZED,
&cache,
&headers
));
}
#[test]
fn reads_only_a_non_empty_membership_type() {
assert_eq!(
membership_type(br#"{"membershipType":"free"}"#).as_deref(),
Some("free")
);
assert_eq!(membership_type(br#"{"membershipType":""}"#), None);
assert_eq!(membership_type(br#"{"subscriptionStatus":"active"}"#), None);
assert_eq!(membership_type(b"not-json"), None);
}
#[test]
fn local_stripe_profile_allows_the_cursor_app_origin() {
let origin = HeaderValue::from_static("vscode-file://vscode-app");
let response = local_stripe_profile(Some(origin.clone())).unwrap();
assert_eq!(
response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
Some(&origin)
);
assert_eq!(
response
.headers()
.get(header::ACCESS_CONTROL_ALLOW_CREDENTIALS),
Some(&HeaderValue::from_static("true"))
);
assert_eq!(
response.headers().get(header::VARY),
Some(&HeaderValue::from_static("Origin"))
);
assert_eq!(
response.headers().get(header::CONTENT_TYPE),
Some(&HeaderValue::from_static("application/json"))
);
}
}
+73 -65
View File
@@ -20,6 +20,9 @@ use crate::{
type BlobSetSender = oneshot::Sender<Result<()>>;
const SET_TIMEOUT: Duration = Duration::from_secs(30 * 60);
const GET_TIMEOUT: Duration = Duration::from_secs(10 * 60);
#[derive(Clone)]
pub struct BlobSynchronizer {
inner: Arc<Inner>,
@@ -73,22 +76,20 @@ impl BlobSynchronizer {
let id = self.inner.store.put_blob(data, edges).await?;
let result = self.ensure_set(&id, data).await;
if let Some(trace) = self.inner.handle.trace() {
trace
.linked_blob(
"blob_set",
"byok_server",
&id,
serde_json::json!({
"byte_count": data.len(),
"status": if result.is_ok() { "acknowledged" } else { "error" },
"error": result.as_ref().err().map(ToString::to_string),
"edges": edges.iter().map(|edge| serde_json::json!({
"child_blob_id": edge.child.to_base64(),
"field_name": edge.field_name,
})).collect::<Vec<_>>(),
}),
)
.await;
trace.linked_blob(
"blob_set",
"byok_server",
&id,
serde_json::json!({
"byte_count": data.len(),
"status": if result.is_ok() { "acknowledged" } else { "error" },
"error": result.as_ref().err().map(ToString::to_string),
"edges": edges.iter().map(|edge| serde_json::json!({
"child_blob_id": edge.child.to_base64(),
"field_name": edge.field_name,
})).collect::<Vec<_>>(),
}),
);
}
result?;
Ok(id)
@@ -130,7 +131,7 @@ impl BlobSynchronizer {
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
_ = tokio::time::sleep(SET_TIMEOUT) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.set_requests.lock().await.remove(&id);
@@ -141,18 +142,16 @@ impl BlobSynchronizer {
pub async fn get(&self, blob_id: &BlobId) -> Result<Option<Vec<u8>>> {
if let Some(data) = self.inner.store.get_blob(blob_id).await? {
if let Some(trace) = self.inner.handle.trace() {
trace
.linked_blob(
"blob_get",
"byok_server",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "local_store",
"status": "found",
}),
)
.await;
trace.linked_blob(
"blob_get",
"byok_server",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "local_store",
"status": "found",
}),
);
}
return Ok(Some(data));
}
@@ -183,7 +182,7 @@ impl BlobSynchronizer {
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
_ = tokio::time::sleep(GET_TIMEOUT) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.get_requests.lock().await.remove(&id);
@@ -191,45 +190,39 @@ impl BlobSynchronizer {
if let Some(trace) = self.inner.handle.trace() {
match &result {
Ok(Some(data)) => {
trace
.linked_blob(
"blob_get",
"cursor_client",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "cursor_client",
"status": "found",
}),
)
.await;
trace.linked_blob(
"blob_get",
"cursor_client",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "cursor_client",
"status": "found",
}),
);
}
Ok(None) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "missing",
}),
)
.await;
trace.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "missing",
}),
);
}
Err(error) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "error",
"error": error.to_string(),
}),
)
.await;
trace.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "error",
"error": error.to_string(),
}),
);
}
}
}
@@ -318,3 +311,18 @@ impl BlobSynchronizer {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn set_timeout_allows_slow_cursor_acknowledgements() {
assert_eq!(SET_TIMEOUT, Duration::from_secs(30 * 60));
}
#[test]
fn get_timeout_allows_slow_cursor_responses() {
assert_eq!(GET_TIMEOUT, Duration::from_secs(10 * 60));
}
}
@@ -0,0 +1,411 @@
//! Generates Git commit messages locally for Cursor's SCM action.
//!
//! Cursor sends `aiserver.v1.AiService/WriteGitCommitMessage` with the staged
//! diffs. Empty commit-settings `model_id` keeps the original behaviour and
//! forwards the RPC unchanged (直连). A configured local model identifier
//! answers the request locally: truncated diffs + previous commits form the user
//! message, the customizable commit prompt is the system prompt, and the raw
//! completion is cleaned before being returned.
use std::{
sync::Arc,
time::{Duration, Instant},
};
use axum::{
body::{to_bytes, Body},
extract::{Extension, State},
http::{header, HeaderValue, Request, Response, StatusCode},
};
use futures_util::StreamExt;
use prost::Message;
use tokio_util::sync::CancellationToken;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{
protocol::{connect, proto::aiserver::v1 as ai},
transport::TransportRegistry,
},
model::{
ContentPart, ModelInvocation, ModelRequest, ModelSpec, ProjectedContent, ProjectedMessage,
PromptSpec, Role,
},
plugin::ADAPTER_ID_PREFIX,
provider::{ModelEvent, Provider},
store::CommitSettings,
Error, Result,
};
const DIFF_TOTAL_LIMIT: usize = 40_000;
const DIFF_SINGLE_LIMIT: usize = 16_000;
const PREVIOUS_COMMIT_LIMIT: usize = 12;
const EXPLICIT_CONTEXT_LIMIT: usize = 20_000;
const GENERATION_TIMEOUT: Duration = Duration::from_secs(180);
pub async fn write_git_commit_message(
State(registry): State<TransportRegistry>,
Extension(upstream): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let settings = registry.store().commit_settings().await?;
if settings.is_direct() {
return forward_direct(&registry, upstream, request).await;
}
generate_local(&registry, request, settings).await
}
async fn forward_direct(
registry: &TransportRegistry,
upstream: CursorProxy,
request: Request<Body>,
) -> Result<Response<Body>> {
let settings = registry.store().tab_settings().await?;
match settings.service_url() {
Some(service_url) => proxy::forward_to_service(&upstream, request, service_url).await,
None => proxy::forward(Extension(upstream), request).await,
}
}
async fn generate_local(
registry: &TransportRegistry,
request: Request<Body>,
settings: CommitSettings,
) -> Result<Response<Body>> {
let (parts, body) = request.into_parts();
let connect_timeout_ms = parts
.headers
.get("connect-timeout-ms")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok());
tracing::info!(?connect_timeout_ms, "write git commit message received");
let body = to_bytes(body, usize::MAX)
.await
.map_err(|error| Error::Protocol(format!("cannot read request body: {error}")))?;
let request: ai::WriteGitCommitMessageRequest = connect::decode_unary(&body)?;
let diffs = truncate_diffs(&request.diffs, DIFF_TOTAL_LIMIT, DIFF_SINGLE_LIMIT);
if diffs.is_empty() {
return Err(Error::Protocol("diffs are required".into()));
}
let model_id = settings.model_id.trim();
ensure_configured_model(registry, model_id).await?;
let invocation = build_invocation(&settings, model_id, build_user_content(&request, &diffs));
let provider = registry.conversations().dependencies().provider.clone();
let generated = generate(
provider,
invocation,
connect_timeout_ms.map(Duration::from_millis),
)
.await?;
let commit_message = clean_generated_commit_message(&generated);
if commit_message.is_empty() {
return Err(Error::Provider("generated commit message is empty".into()));
}
let payload = ai::WriteGitCommitMessageResponse { commit_message }.encode_to_vec();
let mut response = Response::new(Body::from(payload));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
async fn ensure_configured_model(registry: &TransportRegistry, model_id: &str) -> Result<()> {
if model_id.starts_with(ADAPTER_ID_PREFIX) {
let plugins = registry.plugins().ok_or_else(|| {
Error::Provider(format!("commit plugin model {model_id} is unavailable"))
})?;
plugins.model_descriptor(model_id).await?;
return Ok(());
}
if registry.store().model(model_id).await?.is_some() {
return Ok(());
}
Err(Error::Provider(format!(
"commit model {model_id} is not configured; select a configured model in the commit settings"
)))
}
fn build_invocation(
settings: &CommitSettings,
model_id: &str,
user_content: String,
) -> ModelInvocation {
let call_id = format!("commit-message-{}", uuid::Uuid::new_v4());
ModelInvocation {
call_id: call_id.clone(),
run_id: call_id.clone(),
conversation_id: call_id,
provider_call_index: 0,
request: ModelRequest {
prompt: PromptSpec {
instructions: settings.effective_prompt().to_owned(),
tools: Vec::new(),
},
model: ModelSpec::new(model_id.to_owned()),
history: vec![ProjectedMessage {
message_id: "commit-message".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text { text: user_content }]),
}],
},
}
}
async fn generate(
provider: Arc<dyn Provider>,
invocation: ModelInvocation,
client_timeout: Option<Duration>,
) -> Result<String> {
let cancellation = CancellationToken::new();
let stream = provider.stream(invocation, cancellation.clone());
let mut accumulated = String::new();
let soft_deadline = client_timeout
.map(|timeout| timeout.saturating_sub(Duration::from_millis(700)))
.filter(|deadline| !deadline.is_zero());
let deadline = Instant::now()
+ soft_deadline
.unwrap_or(GENERATION_TIMEOUT)
.min(GENERATION_TIMEOUT);
let mut deadline_hit = false;
let completed = tokio::time::timeout(GENERATION_TIMEOUT, async {
futures_util::pin_mut!(stream);
let mut finished = false;
loop {
let wait = deadline.saturating_duration_since(Instant::now());
if wait.is_zero() {
deadline_hit = true;
break;
}
match tokio::time::timeout(wait, stream.next()).await {
Err(_) => {
deadline_hit = true;
break;
}
Ok(None) => break,
Ok(Some(event)) => match event? {
ModelEvent::TextDelta(delta) => accumulated.push_str(&delta),
ModelEvent::ToolCallStart { .. } => {
return Err(Error::Provider(
"commit message generation must not invoke tools".into(),
));
}
ModelEvent::Done(_) => {
finished = true;
break;
}
_ => {}
},
}
}
if !finished && !deadline_hit {
return Err(Error::Provider(
"provider stream ended without Done during commit message generation".into(),
));
}
Ok(())
})
.await;
match completed {
Ok(result) => {
result?;
if accumulated.trim().is_empty() {
return Err(Error::Provider("generated commit message is empty".into()));
}
Ok(accumulated)
}
Err(_) => {
cancellation.cancel();
Err(Error::Provider(
"commit message generation timed out".into(),
))
}
}
}
fn build_user_content(request: &ai::WriteGitCommitMessageRequest, diffs: &[String]) -> String {
let mut sections = vec!["Generate a Git commit message for the following changes.".to_owned()];
let previous =
truncate_previous_commits(&request.previous_commit_messages, PREVIOUS_COMMIT_LIMIT);
if !previous.is_empty() {
sections.push(format!("Recent commit messages:\n{}", previous.join("\n")));
}
if let Some(context) = &request.explicit_context {
let context_json = explicit_context_json(context);
if !context_json.is_empty() {
sections.push(format!("Explicit context:\n{context_json}"));
}
}
let diff_sections: Vec<String> = diffs
.iter()
.enumerate()
.map(|(index, diff)| format!("--- Diff {} ---\n{}", index + 1, diff))
.collect();
sections.push(format!("Diffs:\n{}", diff_sections.join("\n\n")));
sections.join("\n\n")
}
fn explicit_context_json(context: &ai::ExplicitContext) -> String {
let context_text = context.context.trim();
let repo_context = context
.repo_context
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty());
if context_text.is_empty() && repo_context.is_none() {
return String::new();
}
let mut fields = serde_json::Map::new();
if !context_text.is_empty() {
fields.insert(
"context".into(),
serde_json::Value::String(context_text.to_owned()),
);
}
if let Some(repo_context) = repo_context {
fields.insert(
"repo_context".into(),
serde_json::Value::String(repo_context.to_owned()),
);
}
truncate_text(
&serde_json::to_string(&serde_json::Value::Object(fields)).unwrap_or_default(),
EXPLICIT_CONTEXT_LIMIT,
)
}
fn truncate_diffs(input: &[String], total_limit: usize, single_limit: usize) -> Vec<String> {
let mut result = Vec::new();
let mut remaining = total_limit;
for raw in input {
let diff = raw.trim();
if diff.is_empty() || remaining == 0 {
continue;
}
let truncated = truncate_text(diff, remaining.min(single_limit));
if truncated.is_empty() {
continue;
}
remaining = remaining.saturating_sub(truncated.chars().count());
result.push(truncated);
}
result
}
fn truncate_previous_commits(input: &[String], limit: usize) -> Vec<String> {
let mut result = Vec::new();
for raw in input {
let value = raw.trim();
if value.is_empty() {
continue;
}
result.push(format!("- {value}"));
if result.len() >= limit {
break;
}
}
result
}
fn truncate_text(value: &str, limit: usize) -> String {
let trimmed = value.trim();
if trimmed.is_empty() || limit == 0 {
return String::new();
}
if trimmed.chars().count() <= limit {
return trimmed.to_owned();
}
let truncated: String = trimmed.chars().take(limit).collect();
format!("{}\n...[truncated]", truncated.trim_end())
}
fn clean_generated_commit_message(value: &str) -> String {
let mut result = strip_code_fence(value.trim());
const PREFIXES: [&str; 3] = ["commit message:", "git commit message:", "message:"];
loop {
let lower = result.trim().to_ascii_lowercase();
let Some(prefix) = PREFIXES.iter().find(|prefix| lower.starts_with(*prefix)) else {
break;
};
result = result.trim()[prefix.len()..].to_owned();
}
let result = result.trim().to_owned();
if result.lines().all(|line| line.trim().is_empty()) {
String::new()
} else {
result
}
}
fn strip_code_fence(value: &str) -> String {
let trimmed = value.trim();
if !trimmed.starts_with("```") {
return trimmed.to_owned();
}
let mut lines = trimmed.lines();
if lines.next().is_none() {
return trimmed.to_owned();
}
let mut body: Vec<&str> = lines.collect();
if body
.last()
.is_some_and(|line| line.trim_start().starts_with("```"))
{
body.pop();
}
body.join("\n").trim().to_owned()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_model_id_is_direct() {
assert!(CommitSettings::default().is_direct());
assert!(CommitSettings {
model_id: " ".into(),
..CommitSettings::default()
}
.is_direct());
assert!(!CommitSettings {
model_id: "abc".into(),
..CommitSettings::default()
}
.is_direct());
}
#[test]
fn truncation_limits_apply_per_diff_and_in_total() {
let diffs = vec!["a".repeat(20_000), "b".repeat(20_000), "c".repeat(20_000)];
let truncated = truncate_diffs(&diffs, DIFF_TOTAL_LIMIT, DIFF_SINGLE_LIMIT);
assert_eq!(truncated.len(), 3);
let total: usize = truncated.iter().map(|diff| diff.chars().count()).sum();
assert!(total <= DIFF_TOTAL_LIMIT + 3 * "\n...[truncated]".len());
}
#[test]
fn cleaning_strips_fences_and_prefixes() {
let raw = "```\nCommit message: fix: 修复登录超时问题\n```";
assert_eq!(clean_generated_commit_message(raw), "fix: 修复登录超时问题");
}
#[test]
fn cleaning_returns_empty_for_blank_output() {
assert_eq!(clean_generated_commit_message(" \n "), "");
}
#[test]
fn explicit_context_drops_empty_fields() {
let empty = explicit_context_json(&ai::ExplicitContext {
context: " ".into(),
repo_context: None,
});
assert_eq!(empty, "");
let filled = explicit_context_json(&ai::ExplicitContext {
context: "背景".into(),
repo_context: Some("repo".into()),
});
assert_eq!(filled, "{\"context\":\"背景\",\"repo_context\":\"repo\"}");
}
}
+119
View File
@@ -0,0 +1,119 @@
//! Routes Cursor metadata calls according to the supplied authentication token.
use axum::{
body::{to_bytes, Body},
extract::{Extension, Request},
http::{header, HeaderValue, Response},
};
use prost::Message;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::protocol::proto::agent::v1 as agent,
local_app, Result,
};
#[derive(Clone, Copy, PartialEq, Message)]
struct EmptyResponse {}
pub async fn available_docs(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn effective_user_plugins(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn user_privacy_mode(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn update_conversation_metadata(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(
&proxy,
request,
agent::UpdateConversationMetadataResponse {},
)
.await
}
async fn route<M: Message>(
proxy: &CursorProxy,
request: Request<Body>,
mock: M,
) -> Result<Response<Body>> {
let local = local_app::request_uses_local_cursor_token(request.headers());
if local {
consume_body(request).await?;
return Ok(proto(mock));
}
proxy::forward(Extension(proxy.clone()), request).await
}
async fn consume_body(request: Request<Body>) -> Result<()> {
to_bytes(request.into_body(), usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok(())
}
fn proto(message: impl Message) -> 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,
HeaderValue::from_static("application/proto"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
HeaderValue::from_str(&length.to_string()).expect("body length is valid"),
);
response
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use axum::body::{to_bytes, Bytes};
use futures_util::stream;
use super::*;
#[tokio::test]
async fn local_response_consumes_request_body_before_replying() {
let polled = Arc::new(AtomicBool::new(false));
let observed = polled.clone();
let body = Body::from_stream(stream::once(async move {
observed.store(true, Ordering::SeqCst);
Ok::<_, std::convert::Infallible>(Bytes::from_static(b"request"))
}));
consume_body(Request::new(body)).await.unwrap();
assert!(polled.load(Ordering::SeqCst));
}
#[tokio::test]
async fn empty_unary_mock_is_bare_protobuf_without_connect_frame() {
let response = proto(EmptyResponse {});
assert_eq!(response.headers()[header::CONTENT_LENGTH], "0");
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
assert!(body.is_empty());
}
}
+143
View File
@@ -0,0 +1,143 @@
//! Caches a recently confirmed official Cursor Free entitlement.
use std::{
sync::Arc,
time::{Duration, Instant},
};
use axum::http::{header, HeaderMap};
use parking_lot::Mutex;
use sha2::{Digest, Sha256};
use crate::local_app;
const FREE_ENTITLEMENT_TTL: Duration = Duration::from_secs(5 * 60);
#[derive(Clone, Default)]
pub struct FreeEntitlementCache {
cached: Arc<Mutex<Option<CachedFreeEntitlement>>>,
}
struct CachedFreeEntitlement {
token_hash: [u8; 32],
expires_at: Instant,
}
impl FreeEntitlementCache {
pub fn is_confirmed_free(&self, headers: &HeaderMap) -> bool {
let Some(token_hash) = official_token_hash(headers) else {
return false;
};
let now = Instant::now();
let mut cached = self.cached.lock();
match cached.as_ref() {
Some(entry) if entry.expires_at > now && entry.token_hash == token_hash => true,
Some(entry) if entry.expires_at <= now => {
*cached = None;
false
}
_ => false,
}
}
pub fn observe_membership(&self, headers: &HeaderMap, membership_type: &str) -> bool {
let Some(token_hash) = official_token_hash(headers) else {
return false;
};
let mut cached = self.cached.lock();
if membership_type.eq_ignore_ascii_case("free") {
*cached = Some(CachedFreeEntitlement {
token_hash,
expires_at: Instant::now() + FREE_ENTITLEMENT_TTL,
});
} else if cached
.as_ref()
.is_some_and(|entry| entry.token_hash == token_hash)
{
*cached = None;
}
true
}
}
fn official_token_hash(headers: &HeaderMap) -> Option<[u8; 32]> {
if local_app::request_uses_local_cursor_token(headers) {
return None;
}
let token = headers
.get(header::AUTHORIZATION)?
.to_str()
.ok()?
.strip_prefix("Bearer ")?;
if token.is_empty() {
return None;
}
Some(Sha256::digest(token.as_bytes()).into())
}
#[cfg(test)]
mod tests {
use super::*;
fn official_headers(token: &str) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
format!("Bearer {token}").parse().unwrap(),
);
headers
}
#[test]
fn caches_only_a_confirmed_free_official_token() {
let cache = FreeEntitlementCache::default();
let free = official_headers("official-free-token");
let other = official_headers("another-token");
cache.observe_membership(&free, "free");
assert!(cache.is_confirmed_free(&free));
assert!(!cache.is_confirmed_free(&other));
}
#[test]
fn confirmed_non_free_membership_clears_the_same_token() {
let cache = FreeEntitlementCache::default();
let headers = official_headers("official-token");
cache.observe_membership(&headers, "free");
cache.observe_membership(&headers, "pro");
assert!(!cache.is_confirmed_free(&headers));
}
#[test]
fn local_token_never_enters_the_entitlement_cache() {
let cache = FreeEntitlementCache::default();
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
crate::local_app::local_cursor_authorization()
.parse()
.unwrap(),
);
assert!(!cache.observe_membership(&headers, "free"));
assert!(!cache.is_confirmed_free(&headers));
assert!(cache.cached.lock().is_none());
}
#[test]
fn expired_confirmation_is_removed() {
let cache = FreeEntitlementCache::default();
let headers = official_headers("official-token");
let token_hash = official_token_hash(&headers).unwrap();
*cache.cached.lock() = Some(CachedFreeEntitlement {
token_hash,
expires_at: Instant::now() - Duration::from_secs(1),
});
assert!(!cache.is_confirmed_free(&headers));
assert!(cache.cached.lock().is_none());
}
}
+4
View File
@@ -3,9 +3,13 @@
pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod commit_message;
pub mod compatibility;
pub mod context_sync;
pub(crate) mod entitlement;
pub mod knowledge;
pub mod model_catalog;
pub mod observability;
pub mod server_config;
pub mod tab;
pub mod usage;
+216
View File
@@ -187,6 +187,32 @@ struct UsableModelsAddition {
models: Vec<agent::ModelDetails>,
}
#[derive(Clone, PartialEq, Message)]
struct DefaultModelResponse {
#[prost(string, tag = "1")]
model: String,
#[prost(string, tag = "2")]
thinking_model: String,
#[prost(bool, tag = "3")]
max_mode: bool,
#[prost(string, tag = "4")]
next_default_set_date: String,
}
#[derive(Clone, PartialEq, Message)]
struct DefaultModelNudgeDataResponse {
#[prost(string, tag = "1")]
nudge_date: String,
#[prost(bool, tag = "2")]
should_default_switch_on_new_chat: bool,
#[prost(string, repeated, tag = "3")]
models_with_no_default_switch: Vec<String>,
#[prost(string, tag = "4")]
conversion_model_override: String,
}
const CLI_LOCAL_MODEL_API_KEY: &str = "cursor-byok-local";
const CONTEXTS: [(&str, &str); 5] = [
("200k", "200K"),
("356k", "356K"),
@@ -287,6 +313,94 @@ pub async fn usable_models(
}
}
pub async fn default_model_for_cli(
State(registry): State<TransportRegistry>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
agent::GetDefaultModelForCliResponse {
model: default_model_details(&models, &plugin_models),
}
.encode_to_vec(),
))
}
pub async fn default_model(State(registry): State<TransportRegistry>) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
default_model_response(&models, &plugin_models).encode_to_vec(),
))
}
pub async fn default_model_nudge(
State(registry): State<TransportRegistry>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
default_model_nudge_response(&models, &plugin_models).encode_to_vec(),
))
}
async fn configured_plugin_models(registry: &TransportRegistry) -> Vec<PluginModelDescriptor> {
match registry.plugins() {
Some(plugins) => plugins.configured_models().await,
None => Vec::new(),
}
}
fn default_model_details(
models: &[ModelConfig],
plugin_models: &[PluginModelDescriptor],
) -> Option<agent::ModelDetails> {
models
.first()
.map(usable_model)
.or_else(|| plugin_models.first().map(usable_plugin_model))
}
fn default_model_id<'a>(
models: &'a [ModelConfig],
plugin_models: &'a [PluginModelDescriptor],
) -> &'a str {
models
.first()
.map(|model| model.model_hash.as_str())
.or_else(|| plugin_models.first().map(|model| model.id.as_str()))
.unwrap_or_default()
}
fn default_model_response(
models: &[ModelConfig],
plugin_models: &[PluginModelDescriptor],
) -> DefaultModelResponse {
let model = default_model_id(models, plugin_models).to_owned();
DefaultModelResponse {
thinking_model: model.clone(),
model,
max_mode: false,
next_default_set_date: String::new(),
}
}
fn default_model_nudge_response(
models: &[ModelConfig],
plugin_models: &[PluginModelDescriptor],
) -> DefaultModelNudgeDataResponse {
DefaultModelNudgeDataResponse {
nudge_date: "0".into(),
should_default_switch_on_new_chat: false,
models_with_no_default_switch: models
.iter()
.map(|model| model.model_hash.clone())
.chain(plugin_models.iter().map(|model| model.id.clone()))
.collect(),
conversion_model_override: String::new(),
}
}
fn merge_response(upstream: proxy::BufferedResponse, extra: Vec<u8>) -> Result<Response<Body>> {
if !upstream.status.is_success() {
tracing::warn!(status = %upstream.status, "Cursor model catalog upstream rejected request; using local catalog");
@@ -622,6 +736,13 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
}
}
fn cli_local_model_credentials() -> agent::model_details::Credentials {
agent::model_details::Credentials::ApiKeyCredentials(agent::ApiKeyCredentials {
api_key: CLI_LOCAL_MODEL_API_KEY.into(),
base_url: None,
})
}
fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
agent::ModelDetails {
model_id: model.id.clone(),
@@ -629,6 +750,7 @@ fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
display_name: model.display_name.clone(),
display_name_short: model.display_name.clone(),
thinking_details: Some(agent::ThinkingDetails::default()),
credentials: Some(cli_local_model_credentials()),
..Default::default()
}
}
@@ -640,6 +762,100 @@ fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
display_name: model.display_name.clone(),
display_name_short: model.display_name.clone(),
thinking_details: Some(agent::ThinkingDetails::default()),
credentials: Some(cli_local_model_credentials()),
..Default::default()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{ModelType, OPENAI_CHAT_ENDPOINT};
fn model() -> ModelConfig {
ModelConfig {
model_hash: "local-model-hash".into(),
sort_order: 0,
display_name: "Local Model".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://provider.example/v1/chat/completions".into(),
use_full_url: true,
api_key: "provider-secret".into(),
tooltip_data: "Local Model".into(),
model_id: "upstream-model".into(),
reasoning_effort: None,
openai_endpoint: 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,
created_at_ms: 0,
updated_at_ms: 0,
}
}
#[test]
fn cli_model_details_use_local_routing_credentials() {
let details = usable_model(&model());
assert_eq!(details.model_id, "local-model-hash");
assert_eq!(details.display_name, "Local Model");
let agent::model_details::Credentials::ApiKeyCredentials(credentials) =
details.credentials.expect("API credentials")
else {
panic!("expected API key credentials");
};
assert_eq!(credentials.api_key, CLI_LOCAL_MODEL_API_KEY);
assert_eq!(credentials.base_url, None);
assert_ne!(credentials.api_key, "provider-secret");
}
#[test]
fn cli_plugin_model_details_use_local_routing_credentials() {
let details = usable_plugin_model(&PluginModelDescriptor {
id: "plugin:test/provider/model".into(),
plugin_id: "plugin:test".into(),
plugin_name: "Test Plugin".into(),
provider_id: "provider".into(),
model_id: "model".into(),
display_name: "Plugin Model".into(),
description: None,
icon: String::new(),
provider_type: "test".into(),
max_output_tokens: None,
images: false,
});
assert_eq!(details.model_id, "plugin:test/provider/model");
let agent::model_details::Credentials::ApiKeyCredentials(credentials) =
details.credentials.expect("API credentials")
else {
panic!("expected API key credentials");
};
assert_eq!(credentials.api_key, CLI_LOCAL_MODEL_API_KEY);
assert_eq!(credentials.base_url, None);
}
#[test]
fn cli_default_responses_use_the_local_model_hash() {
let models = vec![model()];
let details = default_model_details(&models, &[]).expect("default model");
assert_eq!(details.model_id, "local-model-hash");
let response = default_model_response(&models, &[]);
assert_eq!(response.model, "local-model-hash");
assert_eq!(response.thinking_model, "local-model-hash");
let nudge = default_model_nudge_response(&models, &[]);
assert_eq!(
nudge.models_with_no_default_switch,
vec!["local-model-hash"]
);
}
}
-225
View File
@@ -1,225 +0,0 @@
//! Records Cursor request traces and artifacts.
use std::{
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::{Duration, Instant},
};
use tokio::sync::Mutex;
use crate::store::{BlobId, BufferedCursorTraceChunk, Store};
#[derive(Clone)]
pub struct CursorTraceRecorder {
store: Store,
request_id: String,
chunks: Arc<Mutex<TraceChunkBuffer>>,
finished: Arc<AtomicBool>,
}
#[derive(Default)]
struct TraceChunkBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
first_chunk_at: Option<Instant>,
generation: u64,
}
const MAX_BUFFERED_CHUNKS: usize = 32;
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
const MAX_BUFFER_AGE: Duration = Duration::from_millis(50);
impl CursorTraceRecorder {
pub async fn begin(
store: Store,
request_id: &str,
conversation_id: Option<&str>,
route: &str,
model_id: Option<&str>,
) -> Option<Self> {
match store
.start_cursor_trace_if_detailed(request_id, conversation_id, route, model_id)
.await
{
Ok(true) => Some(Self {
store,
request_id: request_id.into(),
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
finished: Arc::new(AtomicBool::new(false)),
}),
Ok(false) => None,
Err(error) => {
tracing::warn!(request_id, %error, "failed to start Cursor trace");
None
}
}
}
pub async fn resume(store: Store, request_id: &str) -> Option<Self> {
match store.cursor_trace_exists(request_id).await {
Ok(true) => Some(Self {
store,
request_id: request_id.into(),
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
finished: Arc::new(AtomicBool::new(false)),
}),
Ok(false) => None,
Err(error) => {
tracing::warn!(request_id, %error, "failed to resume Cursor trace");
None
}
}
}
pub fn request_id(&self) -> &str {
&self.request_id
}
pub async fn request(&self, artifact_type: &str, data: &[u8], metadata: serde_json::Value) {
if let Err(error) = self
.store
.append_cursor_trace_artifact(
&self.request_id,
artifact_type,
"cursor_client",
data,
&metadata,
)
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor request artifact");
return;
}
if let Err(error) = self
.store
.add_cursor_trace_request_bytes(&self.request_id, data.len())
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to update Cursor request trace size");
}
}
pub async fn artifact(
&self,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: serde_json::Value,
) {
if let Err(error) = self
.store
.append_cursor_trace_artifact(&self.request_id, artifact_type, source, data, &metadata)
.await
{
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to record Cursor trace artifact");
}
}
pub async fn linked_blob(
&self,
artifact_type: &str,
source: &str,
blob_id: &BlobId,
metadata: serde_json::Value,
) {
if let Err(error) = self
.store
.link_cursor_trace_artifact(&self.request_id, artifact_type, source, blob_id, &metadata)
.await
{
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to link Cursor trace Blob");
}
}
pub async fn response_started(&self, status: u16) {
if let Err(error) = self
.store
.start_cursor_trace_response(&self.request_id, status)
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to start Cursor response trace");
}
}
pub async fn response_chunk(&self, source: &str, data: &[u8]) {
let mut buffer = self.chunks.lock().await;
if self.finished.load(Ordering::Acquire) {
return;
}
let schedule_flush = if buffer.chunks.is_empty() {
buffer.generation = buffer.generation.wrapping_add(1);
buffer.first_chunk_at = Some(Instant::now());
Some(buffer.generation)
} else {
None
};
buffer.bytes += data.len();
buffer
.chunks
.push(BufferedCursorTraceChunk::new(source, data));
let expired = buffer
.first_chunk_at
.is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE);
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS
|| buffer.bytes >= MAX_BUFFERED_BYTES
|| expired
{
if let Err(error) = self.flush_locked(&mut buffer).await {
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk");
}
}
drop(buffer);
if let Some(generation) = schedule_flush {
let recorder = self.clone();
tokio::spawn(async move {
tokio::time::sleep(MAX_BUFFER_AGE).await;
let mut buffer = recorder.chunks.lock().await;
if buffer.generation == generation {
if let Err(error) = recorder.flush_locked(&mut buffer).await {
tracing::warn!(request_id = recorder.request_id, %error, "failed to flush Cursor response chunks");
}
}
});
}
}
pub async fn finish(&self, error: Option<&str>) {
if self.finished.swap(true, Ordering::AcqRel) {
return;
}
let mut buffer = self.chunks.lock().await;
if let Err(store_error) = self.flush_locked(&mut buffer).await {
tracing::warn!(request_id = self.request_id, %store_error, "failed to flush Cursor response chunks");
}
drop(buffer);
if let Err(store_error) = self
.store
.finish_cursor_trace(&self.request_id, error)
.await
{
tracing::warn!(request_id = self.request_id, %store_error, "failed to finish Cursor trace");
}
}
async fn flush_locked(&self, buffer: &mut TraceChunkBuffer) -> crate::Result<()> {
if buffer.chunks.is_empty() {
return Ok(());
}
let chunks = std::mem::take(&mut buffer.chunks);
buffer.bytes = 0;
buffer.first_chunk_at = None;
if let Err(error) = self
.store
.add_cursor_trace_response_chunks(&self.request_id, &chunks)
.await
{
buffer.bytes = chunks.iter().map(|chunk| chunk.data.len()).sum();
buffer.first_chunk_at = Some(Instant::now());
buffer.chunks = chunks;
return Err(error);
}
Ok(())
}
}
@@ -0,0 +1,71 @@
use std::sync::{atomic::AtomicU8, Arc};
use bytes::Bytes;
use crate::store::BlobId;
pub(super) const TRACE_UNKNOWN: u8 = 0;
pub(super) const TRACE_ACTIVE: u8 = 1;
pub(super) const TRACE_DISABLED: u8 = 2;
pub(super) enum TraceEvent {
Begin {
request_id: String,
activation: Arc<AtomicU8>,
conversation_id: Option<String>,
route: String,
model_id: Option<String>,
},
Resume {
request_id: String,
activation: Arc<AtomicU8>,
},
Request {
request_id: String,
artifact_type: String,
data: Bytes,
metadata: serde_json::Value,
},
Artifact {
request_id: String,
artifact_type: String,
source: String,
data: Bytes,
metadata: serde_json::Value,
},
LinkedBlob {
request_id: String,
artifact_type: String,
source: String,
blob_id: BlobId,
metadata: serde_json::Value,
},
ResponseStarted {
request_id: String,
status: u16,
},
ResponseChunk {
request_id: String,
source: String,
data: Bytes,
},
Finish {
request_id: String,
error: Option<String>,
},
}
impl TraceEvent {
pub(super) fn request_id(&self) -> &str {
match self {
Self::Begin { request_id, .. }
| Self::Resume { request_id, .. }
| Self::Request { request_id, .. }
| Self::Artifact { request_id, .. }
| Self::LinkedBlob { request_id, .. }
| Self::ResponseStarted { request_id, .. }
| Self::ResponseChunk { request_id, .. }
| Self::Finish { request_id, .. } => request_id,
}
}
}
@@ -0,0 +1,157 @@
//! Records Cursor request traces without blocking request or runtime paths.
mod event;
mod worker;
use std::sync::{
atomic::{AtomicBool, AtomicU8, Ordering},
Arc,
};
use bytes::Bytes;
use tokio::sync::mpsc;
use crate::store::{BlobId, Store};
use event::{TraceEvent, TRACE_DISABLED, TRACE_UNKNOWN};
const TRACE_QUEUE_CAPACITY: usize = 512;
#[derive(Clone)]
pub struct CursorTraceService {
sender: mpsc::Sender<TraceEvent>,
}
impl CursorTraceService {
pub fn new(store: Store) -> Self {
let (sender, receiver) = mpsc::channel(TRACE_QUEUE_CAPACITY);
tokio::spawn(worker::run(store, receiver));
Self { sender }
}
pub fn recorder(&self, request_id: &str) -> CursorTraceRecorder {
CursorTraceRecorder {
request_id: Arc::from(request_id),
sender: self.sender.clone(),
finished: Arc::new(AtomicBool::new(false)),
activation: Arc::new(AtomicU8::new(TRACE_UNKNOWN)),
}
}
}
#[derive(Clone)]
pub struct CursorTraceRecorder {
request_id: Arc<str>,
sender: mpsc::Sender<TraceEvent>,
finished: Arc<AtomicBool>,
activation: Arc<AtomicU8>,
}
impl CursorTraceRecorder {
pub fn request_id(&self) -> &str {
&self.request_id
}
pub fn begin(&self, conversation_id: Option<&str>, route: &str, model_id: Option<&str>) {
self.send_control(TraceEvent::Begin {
request_id: self.request_id.to_string(),
activation: self.activation.clone(),
conversation_id: conversation_id.map(str::to_owned),
route: route.to_owned(),
model_id: model_id.map(str::to_owned),
});
}
pub fn resume(&self) {
self.send_control(TraceEvent::Resume {
request_id: self.request_id.to_string(),
activation: self.activation.clone(),
});
}
pub fn request(&self, artifact_type: &str, data: Bytes, metadata: serde_json::Value) {
self.send(TraceEvent::Request {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
data,
metadata,
});
}
pub fn artifact(
&self,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: serde_json::Value,
) {
self.send(TraceEvent::Artifact {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
source: source.to_owned(),
data: Bytes::copy_from_slice(data),
metadata,
});
}
pub fn linked_blob(
&self,
artifact_type: &str,
source: &str,
blob_id: &BlobId,
metadata: serde_json::Value,
) {
self.send(TraceEvent::LinkedBlob {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
source: source.to_owned(),
blob_id: blob_id.clone(),
metadata,
});
}
pub fn response_started(&self, status: u16) {
self.send(TraceEvent::ResponseStarted {
request_id: self.request_id.to_string(),
status,
});
}
pub fn response_chunk(&self, source: &str, data: Bytes) {
if self.finished.load(Ordering::Acquire) {
return;
}
self.send(TraceEvent::ResponseChunk {
request_id: self.request_id.to_string(),
source: source.to_owned(),
data,
});
}
pub fn finish(&self, error: Option<&str>) {
if self.finished.swap(true, Ordering::AcqRel) {
return;
}
self.send_control(TraceEvent::Finish {
request_id: self.request_id.to_string(),
error: error.map(str::to_owned),
});
}
fn send(&self, event: TraceEvent) {
if self.activation.load(Ordering::Acquire) == TRACE_DISABLED {
return;
}
self.send_control(event);
}
fn send_control(&self, event: TraceEvent) {
if let Err(error) = self.sender.try_send(event) {
tracing::warn!(
request_id = %self.request_id,
%error,
"dropping Cursor trace event"
);
}
}
}
@@ -0,0 +1,349 @@
use std::{
collections::{BTreeMap, HashMap},
sync::atomic::Ordering,
time::Duration,
};
use tokio::sync::mpsc;
use crate::store::{BufferedCursorTraceChunk, Store};
use super::event::{TraceEvent, TRACE_ACTIVE, TRACE_DISABLED};
const MAX_BUFFERED_CHUNKS: usize = 32;
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
const FLUSH_INTERVAL: Duration = Duration::from_millis(50);
#[derive(Clone, Copy)]
enum TraceState {
Active,
Disabled,
}
#[derive(Default)]
struct ResponseBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
}
struct BufferedRequest {
artifact_type: String,
data: bytes::Bytes,
metadata: serde_json::Value,
}
#[derive(Default)]
struct RequestOrder {
next: i64,
pending: BTreeMap<i64, Vec<BufferedRequest>>,
}
pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver<TraceEvent>) {
let mut states = HashMap::<String, TraceState>::new();
let mut buffers = HashMap::<String, ResponseBuffer>::new();
let mut request_orders = HashMap::<String, RequestOrder>::new();
let mut interval = tokio::time::interval(FLUSH_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
event = receiver.recv() => {
let Some(event) = event else {
flush_all(&store, &mut buffers).await;
return;
};
process(&store, &mut states, &mut buffers, &mut request_orders, event).await;
}
_ = interval.tick() => flush_all(&store, &mut buffers).await,
}
}
}
async fn process(
store: &Store,
states: &mut HashMap<String, TraceState>,
buffers: &mut HashMap<String, ResponseBuffer>,
request_orders: &mut HashMap<String, RequestOrder>,
event: TraceEvent,
) {
let request_id = event.request_id().to_owned();
let finishes_trace = matches!(&event, TraceEvent::Finish { .. });
match event {
TraceEvent::Begin {
request_id,
activation,
conversation_id,
route,
model_id,
} => {
let state = match store
.start_cursor_trace_if_detailed(
&request_id,
conversation_id.as_deref(),
&route,
model_id.as_deref(),
)
.await
{
Ok(true) => TraceState::Active,
Ok(false) => TraceState::Disabled,
Err(error) => {
tracing::warn!(%request_id, %error, "failed to start Cursor trace");
TraceState::Disabled
}
};
activation.store(
match state {
TraceState::Active => TRACE_ACTIVE,
TraceState::Disabled => TRACE_DISABLED,
},
Ordering::Release,
);
states.insert(request_id, state);
return;
}
TraceEvent::Resume {
request_id,
activation,
} => {
let state = ensure_state(store, states, &request_id).await;
activation.store(
match state {
TraceState::Active => TRACE_ACTIVE,
TraceState::Disabled => TRACE_DISABLED,
},
Ordering::Release,
);
return;
}
_ => {}
}
if !matches!(
ensure_state(store, states, &request_id).await,
TraceState::Active
) {
if finishes_trace {
states.remove(&request_id);
buffers.remove(&request_id);
request_orders.remove(&request_id);
}
return;
}
let result = match event {
TraceEvent::Request {
artifact_type,
data,
metadata,
..
} => {
append_request(
store,
request_orders,
&request_id,
artifact_type,
data,
metadata,
)
.await
}
TraceEvent::Artifact {
artifact_type,
source,
data,
metadata,
..
} => {
store
.append_cursor_trace_artifact(
&request_id,
&artifact_type,
&source,
&data,
&metadata,
)
.await
}
TraceEvent::LinkedBlob {
artifact_type,
source,
blob_id,
metadata,
..
} => {
store
.link_cursor_trace_artifact(
&request_id,
&artifact_type,
&source,
&blob_id,
&metadata,
)
.await
}
TraceEvent::ResponseStarted { status, .. } => {
store.start_cursor_trace_response(&request_id, status).await
}
TraceEvent::ResponseChunk { source, data, .. } => {
let buffer = buffers.entry(request_id.clone()).or_default();
buffer.bytes += data.len();
buffer
.chunks
.push(BufferedCursorTraceChunk::new(&source, &data));
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS || buffer.bytes >= MAX_BUFFERED_BYTES {
flush_one(store, buffers, &request_id).await;
}
return;
}
TraceEvent::Finish { error, .. } => {
flush_request_order(store, request_orders, &request_id).await;
flush_one(store, buffers, &request_id).await;
store
.finish_cursor_trace(&request_id, error.as_deref())
.await
}
TraceEvent::Begin { .. } | TraceEvent::Resume { .. } => unreachable!(),
};
if let Err(error) = result {
tracing::warn!(%request_id, %error, "failed to record Cursor trace event");
}
if finishes_trace {
states.remove(&request_id);
buffers.remove(&request_id);
}
}
async fn append_request(
store: &Store,
request_orders: &mut HashMap<String, RequestOrder>,
request_id: &str,
artifact_type: String,
data: bytes::Bytes,
metadata: serde_json::Value,
) -> crate::Result<()> {
let append_seqno = metadata
.get("append_seqno")
.and_then(serde_json::Value::as_i64);
let ordered = artifact_type == "bidi_request"
&& metadata
.get("accepted")
.and_then(serde_json::Value::as_bool)
== Some(true)
&& metadata
.get("route_outcome")
.and_then(serde_json::Value::as_str)
== Some("local");
let Some(append_seqno) = append_seqno.filter(|_| ordered) else {
return store
.append_cursor_trace_request(
request_id,
&artifact_type,
"cursor_client",
&data,
&metadata,
)
.await;
};
let request = BufferedRequest {
artifact_type,
data,
metadata,
};
let order = request_orders.entry(request_id.to_owned()).or_default();
if append_seqno < order.next {
return store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await;
}
order.pending.entry(append_seqno).or_default().push(request);
while let Some(requests) = order.pending.remove(&order.next) {
for request in requests {
store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await?;
}
order.next = order.next.saturating_add(1);
}
Ok(())
}
async fn flush_request_order(
store: &Store,
request_orders: &mut HashMap<String, RequestOrder>,
request_id: &str,
) {
let Some(order) = request_orders.remove(request_id) else {
return;
};
for requests in order.pending.into_values() {
for request in requests {
if let Err(error) = store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await
{
tracing::warn!(%request_id, %error, "failed to flush ordered Cursor request trace");
}
}
}
}
async fn ensure_state(
store: &Store,
states: &mut HashMap<String, TraceState>,
request_id: &str,
) -> TraceState {
if let Some(state) = states.get(request_id).copied() {
return state;
}
let state = match store.cursor_trace_exists(request_id).await {
Ok(true) => TraceState::Active,
Ok(false) => TraceState::Disabled,
Err(error) => {
tracing::warn!(%request_id, %error, "failed to resume Cursor trace");
TraceState::Disabled
}
};
states.insert(request_id.to_owned(), state);
state
}
async fn flush_one(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>, request_id: &str) {
let Some(mut buffer) = buffers.remove(request_id) else {
return;
};
if let Err(error) = store
.add_cursor_trace_response_chunks(request_id, &buffer.chunks)
.await
{
tracing::warn!(%request_id, %error, "failed to flush Cursor response chunks");
buffer.bytes = buffer.chunks.iter().map(|chunk| chunk.data.len()).sum();
buffers.insert(request_id.to_owned(), buffer);
}
}
async fn flush_all(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>) {
let request_ids = buffers.keys().cloned().collect::<Vec<_>>();
for request_id in request_ids {
flush_one(store, buffers, &request_id).await;
}
}
@@ -0,0 +1,51 @@
//! Keeps Cursor Agent traffic on the endpoint selected with `agent -e`.
use axum::{
body::Body,
http::{header, HeaderValue, Response, StatusCode},
};
use prost::Message;
use crate::Result;
const HTTP2_CONFIG_FORCE_ALL_DISABLED: i32 = 1;
#[derive(Clone, PartialEq, Message)]
struct ServerConfigResponse {
#[prost(string, tag = "6")]
config_version: String,
#[prost(int32, tag = "7")]
http2_config: i32,
#[prost(bool, optional, tag = "28")]
cli_sandbox_default_enabled: Option<bool>,
}
pub async fn get() -> Result<Response<Body>> {
let payload = server_config().encode_to_vec();
let mut response = Response::new(Body::from(payload));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
fn server_config() -> ServerConfigResponse {
ServerConfigResponse {
config_version: "cursor_byok_local_agent_v1".into(),
http2_config: HTTP2_CONFIG_FORCE_ALL_DISABLED,
cli_sandbox_default_enabled: Some(true),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn forces_agent_cli_to_use_the_selected_legacy_endpoint() {
let config = server_config();
assert_eq!(config.http2_config, HTTP2_CONFIG_FORCE_ALL_DISABLED);
assert_eq!(config.cli_sandbox_default_enabled, Some(true));
}
}
+1 -2
View File
@@ -9,7 +9,7 @@ use axum::{
use crate::{api::cursor::proxy, cursor::transport::TransportRegistry, Result};
pub const TAB_PATHS: [&str; 17] = [
pub const TAB_PATHS: [&str; 16] = [
"/aiserver.v1.AiService/StreamCpp",
"/aiserver.v1.AiService/StreamNextCursorPrediction",
"/aiserver.v1.AiService/GetCppEditClassification",
@@ -19,7 +19,6 @@ pub const TAB_PATHS: [&str; 17] = [
"/aiserver.v1.AiService/CppAppend",
"/aiserver.v1.AiService/CppEditHistoryAppend",
"/aiserver.v1.AiService/ReportAiCodeChangeMetrics",
"/aiserver.v1.AiService/WriteGitCommitMessage",
"/aiserver.v1.AiService/WriteGitBranchName",
"/aiserver.v1.CppService/AvailableModels",
"/aiserver.v1.CppService/RecordCppFate",
@@ -12,7 +12,7 @@ pub(super) fn output(
Message::WriteResult(value) => write(value),
Message::DeleteResult(value) => delete(value),
Message::GrepResult(value) => grep(value),
Message::DiagnosticsResult(value) => diagnostics(value),
Message::DiagnosticsResult(value) => diagnostics(value, call),
Message::McpResult(value) => mcp(value),
Message::ReadMcpResourceExecResult(value) => read_mcp(value),
Message::SubagentResult(value) => task(value, call),
@@ -212,14 +212,14 @@ fn grep_truncation(
}
}
fn diagnostics(value: &pb::DiagnosticsResult) -> Result<(String, bool)> {
fn diagnostics(value: &pb::DiagnosticsResult, call: &ToolCall) -> Result<(String, bool)> {
use pb::diagnostics_result::Result as R;
match value
.result
.as_ref()
.ok_or_else(|| missing("diagnostics"))?
{
R::Success(value) => Ok((diagnostics_success(value), false)),
R::Success(value) => Ok((diagnostics_success(value, call), false)),
R::Error(value) => Ok((value.error.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::FileNotFound(value) => Ok((format!("file not found: {}", value.path), true)),
@@ -227,9 +227,40 @@ fn diagnostics(value: &pb::DiagnosticsResult) -> Result<(String, bool)> {
}
}
fn diagnostics_success(value: &pb::DiagnosticsSuccess) -> String {
/// `codec::request` encodes only `paths[0]` into the `DiagnosticsArgs` exec, and
/// `DiagnosticsSuccess` carries a single `path`, so a multi-path ReadLints call
/// only ever inspects the first entry. Name the rest instead of letting a clean
/// result for one file read as a clean bill of health for all of them.
fn unchecked_lint_paths(call: &ToolCall) -> Vec<&str> {
call.arguments
.get("paths")
.and_then(serde_json::Value::as_array)
.map(|paths| {
paths
.iter()
.skip(1)
.filter_map(serde_json::Value::as_str)
.filter(|path| !path.is_empty())
.collect()
})
.unwrap_or_default()
}
fn diagnostics_success(value: &pb::DiagnosticsSuccess, call: &ToolCall) -> String {
let unchecked = unchecked_lint_paths(call);
let notice = (!unchecked.is_empty()).then(|| {
format!(
"[Only {} was checked; ReadLints reads one path per call. Not checked: {}]",
value.path,
unchecked.join(", ")
)
});
if value.diagnostics.is_empty() {
return format!("No diagnostics found in {}", value.path);
let clean = format!("No diagnostics found in {}", value.path);
return match notice {
Some(notice) => format!("{clean}\n{notice}"),
None => clean,
};
}
let mut lines = value
.diagnostics
@@ -261,6 +292,7 @@ fn diagnostics_success(value: &pb::DiagnosticsSuccess) -> String {
value.diagnostics.len()
));
}
lines.extend(notice);
lines.join("\n")
}
@@ -413,3 +445,83 @@ fn creates_subagent(call: &ToolCall) -> bool {
fn missing(name: &str) -> Error {
Error::Protocol(format!("{name} returned no result"))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn read_lints(paths: serde_json::Value) -> ToolCall {
ToolCall {
index: 0,
call_id: "lints".into(),
model_call_id: "model:0".into(),
name: "ReadLints".into(),
arguments_text: String::new(),
arguments: json!({ "paths": paths }),
argument_error: None,
}
}
fn clean_result(path: &str) -> pb::exec_client_message::Message {
pb::exec_client_message::Message::DiagnosticsResult(pb::DiagnosticsResult {
result: Some(pb::diagnostics_result::Result::Success(
pb::DiagnosticsSuccess {
path: path.into(),
diagnostics: Vec::new(),
total_diagnostics: 0,
},
)),
})
}
/// The codec encodes only `paths[0]` into the DiagnosticsArgs exec, so a
/// clean result for that one path must not read as "these files are all
/// clean" for the paths that were never looked at.
#[test]
fn read_lints_names_the_paths_it_did_not_check() {
let call = read_lints(json!(["a.ts", "b.ts", "c.ts"]));
let (content, is_error) = output(&clean_result("a.ts"), &call).unwrap();
assert!(!is_error);
assert!(content.contains("a.ts"));
assert!(
content.contains("b.ts") && content.contains("c.ts"),
"unchecked paths must be reported, got: {content}"
);
}
/// The overwhelmingly common single-path call must be byte-for-byte
/// unchanged.
#[test]
fn read_lints_single_path_result_is_unchanged() {
let call = read_lints(json!(["a.ts"]));
let (content, _) = output(&clean_result("a.ts"), &call).unwrap();
assert_eq!(content, "No diagnostics found in a.ts");
}
/// The notice belongs on a result that did report diagnostics too: those
/// diagnostics are still only a.ts's.
#[test]
fn read_lints_reports_unchecked_paths_alongside_diagnostics() {
let call = read_lints(json!(["a.ts", "b.ts"]));
let message = pb::exec_client_message::Message::DiagnosticsResult(pb::DiagnosticsResult {
result: Some(pb::diagnostics_result::Result::Success(
pb::DiagnosticsSuccess {
path: "a.ts".into(),
diagnostics: vec![pb::Diagnostic {
message: "unused import".into(),
severity: pb::DiagnosticSeverity::Warning as i32,
..Default::default()
}],
total_diagnostics: 1,
},
)),
});
let (content, _) = output(&message, &call).unwrap();
assert!(content.contains("unused import"));
assert!(
content.contains("Not checked: b.ts"),
"unchecked paths must be reported, got: {content}"
);
}
}
+188 -11
View File
@@ -258,7 +258,9 @@ fn gate_grep_content(content: &mut pb::GrepContentResult, budget: &mut GrepBudge
truncated = true;
break;
}
budget.content_bytes -= next_match.content.len();
budget.content_bytes = budget
.content_bytes
.saturating_sub(next_match.content.len());
budget.matches -= 1;
next.matches.push(next_match);
}
@@ -385,7 +387,15 @@ fn gate_mcp(tool: &mut pb::McpToolCall) {
"[truncated: MCP content items exceeded {MCP_CONTENT_ITEM_LIMIT} items; showing {MCP_CONTENT_ITEM_LIMIT} of {original_items} items]"
));
}
let original_text_bytes = success.content.iter().fold(0usize, |total, item| {
let bytes = match item.content.as_ref() {
Some(pb::mcp_tool_result_content_item::Content::Text(text)) => text.text.len(),
_ => 0,
};
total.saturating_add(bytes)
});
let mut remaining_text = MCP_TEXT_LIMIT;
let mut text_truncated = false;
let mut content = Vec::with_capacity(success.content.len() + notices.len());
for mut item in std::mem::take(&mut success.content) {
// MCP images are sent to the client as inline binary data. Truncating
@@ -397,19 +407,23 @@ fn gate_mcp(tool: &mut pb::McpToolCall) {
let original = text.text.clone();
let next = truncate_text("MCP content item", &original, MCP_TEXT_LIMIT);
if remaining_text == 0 {
notices.push(truncation_notice(
"MCP text",
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT.saturating_add(original.len()),
));
text_truncated |= !next.is_empty();
continue;
}
text.text = truncate_text("MCP text", &next, remaining_text);
text_truncated |= text.text != next;
remaining_text = remaining_text.saturating_sub(text.text.len());
}
content.push(item);
}
if text_truncated {
notices.push(truncation_notice(
"MCP text",
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT.saturating_sub(remaining_text),
original_text_bytes,
));
}
content.extend(notices.into_iter().map(mcp_notice));
success.content = content;
}
@@ -630,19 +644,52 @@ fn truncate_text(tool_name: &str, content: &str, limit: usize) -> String {
}
let original = content.len();
let mut shown = limit;
let mut previous = None;
loop {
let notice = format!(
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; showing {shown} of {original} bytes]"
);
let available = limit.saturating_sub(notice.len());
let kept = utf8_prefix(content, available);
if kept.len() == shown {
return format!("{}{notice}", kept.trim_end_matches('\n'));
if notice.len() >= limit {
// The notice alone would blow the budget; keep a plain prefix so the
// result never costs more than `limit` bytes.
return utf8_prefix(content, limit).to_string();
}
let kept = utf8_prefix(content, limit - notice.len());
// `notice.len()` grows with the digit count of `shown`, so `kept.len()`
// can alternate between two values across a power-of-ten boundary.
if kept.len() == shown || previous == Some(kept.len()) {
return truncated_prefix_with_notice(tool_name, content, limit, original, kept.len());
}
previous = Some(shown);
shown = kept.len();
}
}
fn truncated_prefix_with_notice(
tool_name: &str,
content: &str,
limit: usize,
original: usize,
initial_cap: usize,
) -> String {
let mut cap = initial_cap;
loop {
let kept = utf8_prefix(content, cap).trim_end_matches('\n');
let notice = format!(
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; showing {} of {original} bytes]",
kept.len()
);
if notice.len() >= limit {
return utf8_prefix(content, limit).to_string();
}
let available = limit - notice.len();
if kept.len() <= available {
return format!("{kept}{notice}");
}
cap = available;
}
}
fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
if content.len() <= limit {
return content.to_string();
@@ -685,3 +732,133 @@ fn utf8_suffix(value: &str, limit: usize) -> &str {
}
&value[start..]
}
#[cfg(test)]
mod tests {
use super::*;
fn grep_tool(matches: Vec<String>) -> pb::tool_call::Tool {
pb::tool_call::Tool::GrepToolCall(pb::GrepToolCall {
args: None,
result: Some(pb::GrepResult {
result: Some(pb::grep_result::Result::Success(pb::GrepSuccess {
active_editor_result: Some(pb::GrepUnionResult {
result: Some(pb::grep_union_result::Result::Content(
pb::GrepContentResult {
matches: vec![pb::GrepFileMatch {
file: "src/lib.rs".into(),
matches: matches
.into_iter()
.enumerate()
.map(|(index, content)| pb::GrepContentMatch {
line_number: index as i32 + 1,
content,
..Default::default()
})
.collect(),
}],
..Default::default()
},
)),
}),
..Default::default()
})),
}),
})
}
fn mcp_tool(texts: Vec<String>) -> pb::McpToolCall {
pb::McpToolCall {
args: None,
result: Some(pb::McpToolResult {
result: Some(pb::mcp_tool_result::Result::Success(pb::McpSuccess {
content: texts
.into_iter()
.map(|text| pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text,
output_location: None,
},
)),
})
.collect(),
is_error: false,
structured_content: None,
})),
}),
description: None,
}
}
#[test]
fn truncate_text_never_exceeds_its_limit() {
let content = "b".repeat(200);
for limit in 1..=250 {
let output = truncate_text("Grep", &content, limit);
assert!(
output.len() <= limit,
"limit {limit} produced {} bytes",
output.len()
);
}
}
#[test]
fn truncate_text_terminates_when_the_notice_length_oscillates() {
// `limit` values where the notice grows and shrinks with the digit count
// of the reported byte count, so the fixed point is never reached.
assert!(truncate_text("Grep", &"b".repeat(200), 78).len() <= 78);
assert!(truncate_text("Grep", &"b".repeat(200), 170).len() <= 170);
assert!(truncate_text("MCP text", &"b".repeat(500), 82).len() <= 82);
}
#[test]
fn truncate_text_reports_the_actual_utf8_prefix_size() {
let content = "😀".repeat(1_000);
let output = truncate_text("Grep", &content, 81);
let (kept, notice) = output
.split_once("\n\n[truncated:")
.expect("the truncation notice fits");
assert!(
notice.contains(&format!("showing {} of", kept.len())),
"notice must report the actual UTF-8 prefix size: {output}"
);
assert!(output.len() <= 81);
}
#[test]
fn mcp_total_text_truncation_always_adds_a_notice() {
let mut tool = mcp_tool(vec!["a".repeat(MCP_TEXT_LIMIT - 8), "b".repeat(100)]);
gate_mcp(&mut tool);
let success = match tool.result.unwrap().result.unwrap() {
pb::mcp_tool_result::Result::Success(success) => success,
_ => panic!("expected MCP success"),
};
assert_eq!(success.content.len(), 3);
assert!(is_mcp_notice(success.content.last().unwrap()));
}
#[test]
fn grep_content_gate_survives_a_nearly_exhausted_byte_budget() {
// 16 matches leave 16 bytes of the 32 KiB content budget, which is less
// than the truncation notice for the 17th match.
let mut matches = vec!["a".repeat(2047); 16];
matches.push("b".repeat(100));
let mut tool = grep_tool(matches);
let mut content = String::new();
tool_completion("Grep", &mut tool, &mut content);
}
#[test]
fn grep_content_gate_terminates_on_an_oscillating_remaining_budget() {
// The same path, tuned so the remaining budget lands on a `limit` where
// the truncation notice length oscillates.
let mut matches = vec!["a".repeat(2043); 15];
matches.push("a".repeat(2045));
matches.push("b".repeat(200));
let mut tool = grep_tool(matches);
let mut content = String::new();
tool_completion("Grep", &mut tool, &mut content);
}
}
+43 -7
View File
@@ -15,7 +15,7 @@ use crate::{
Error, Result,
};
use super::OutputHub;
use super::{OutputHub, TransportAdmission, TransportLifecycle};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransportParent {
@@ -30,7 +30,8 @@ pub struct TransportHandle {
output: Arc<OutputHub>,
conversation_id: Arc<OnceLock<String>>,
parent: Arc<OnceLock<TransportParent>>,
trace: Option<CursorTraceRecorder>,
trace: CursorTraceRecorder,
lifecycle: TransportLifecycle,
disconnect: CancellationToken,
}
@@ -39,7 +40,7 @@ impl TransportHandle {
request_id: String,
commands: mpsc::Sender<TransportCommand>,
output: Arc<OutputHub>,
trace: Option<CursorTraceRecorder>,
trace: CursorTraceRecorder,
) -> Self {
Self {
request_id,
@@ -48,6 +49,7 @@ impl TransportHandle {
conversation_id: Arc::new(OnceLock::new()),
parent: Arc::new(OnceLock::new()),
trace,
lifecycle: TransportLifecycle::new(),
disconnect: CancellationToken::new(),
}
}
@@ -127,12 +129,46 @@ impl TransportHandle {
self.output.close()
}
pub(crate) async fn wait_closed(&self) {
self.output.wait_closed().await;
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
Some(&self.trace)
}
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
self.trace.as_ref()
pub(crate) fn accepting_appends(&self) -> bool {
self.lifecycle.is_open()
}
pub(crate) fn admit(&self) -> Result<TransportAdmission> {
self.lifecycle
.admit()
.ok_or_else(|| Error::RunNotFound(self.request_id.clone()))
}
pub(crate) fn begin_close(&self) {
self.lifecycle.begin_close();
}
pub(crate) fn admissions_drained(&self) -> bool {
self.lifecycle.admissions_drained()
}
pub(crate) async fn wait_admissions_drained(&self) {
self.lifecycle.wait_admissions_drained().await;
}
pub(crate) fn mark_draining(&self) {
self.lifecycle.mark_draining();
}
pub(crate) fn reopen(&self) {
self.lifecycle.reopen();
}
pub(crate) fn close_transport(&self) {
self.lifecycle.close();
}
pub(crate) async fn wait_transport_closed(&self) {
self.lifecycle.wait_closed().await;
}
pub(crate) fn disconnect_token(&self) -> CancellationToken {
+170
View File
@@ -0,0 +1,170 @@
//! Coordinates append admission with transport shutdown.
use std::sync::Arc;
use tokio::sync::Notify;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum TransportState {
Open,
Closing,
Draining,
Closed,
}
#[derive(Clone)]
pub(crate) struct TransportLifecycle {
inner: Arc<LifecycleInner>,
}
struct LifecycleInner {
state: parking_lot::Mutex<LifecycleState>,
admissions_drained: Notify,
closed: Notify,
}
struct LifecycleState {
state: TransportState,
admissions: usize,
}
pub(crate) struct TransportAdmission {
inner: Arc<LifecycleInner>,
}
impl TransportLifecycle {
pub(crate) fn new() -> Self {
Self {
inner: Arc::new(LifecycleInner {
state: parking_lot::Mutex::new(LifecycleState {
state: TransportState::Open,
admissions: 0,
}),
admissions_drained: Notify::new(),
closed: Notify::new(),
}),
}
}
pub(crate) fn is_open(&self) -> bool {
self.inner.state.lock().state == TransportState::Open
}
pub(crate) fn admit(&self) -> Option<TransportAdmission> {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state != TransportState::Open {
return None;
}
lifecycle.admissions += 1;
Some(TransportAdmission {
inner: self.inner.clone(),
})
}
pub(crate) fn begin_close(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state != TransportState::Open {
return;
}
lifecycle.state = TransportState::Closing;
let drained = lifecycle.admissions == 0;
drop(lifecycle);
if drained {
self.inner.admissions_drained.notify_waiters();
}
}
pub(crate) fn admissions_drained(&self) -> bool {
self.inner.state.lock().admissions == 0
}
pub(crate) async fn wait_admissions_drained(&self) {
loop {
let notified = self.inner.admissions_drained.notified();
if self.inner.state.lock().admissions == 0 {
return;
}
notified.await;
}
}
pub(crate) fn mark_draining(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state == TransportState::Closing && lifecycle.admissions == 0 {
lifecycle.state = TransportState::Draining;
}
}
pub(crate) fn reopen(&self) {
let mut lifecycle = self.inner.state.lock();
if matches!(
lifecycle.state,
TransportState::Closing | TransportState::Draining
) {
lifecycle.state = TransportState::Open;
}
}
pub(crate) fn close(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state == TransportState::Closed {
return;
}
lifecycle.state = TransportState::Closed;
drop(lifecycle);
self.inner.closed.notify_waiters();
}
pub(crate) async fn wait_closed(&self) {
loop {
let notified = self.inner.closed.notified();
if self.inner.state.lock().state == TransportState::Closed {
return;
}
notified.await;
}
}
}
impl Drop for TransportAdmission {
fn drop(&mut self) {
let mut lifecycle = self.inner.state.lock();
lifecycle.admissions = lifecycle.admissions.saturating_sub(1);
let drained = lifecycle.state == TransportState::Closing && lifecycle.admissions == 0;
drop(lifecycle);
if drained {
self.inner.admissions_drained.notify_waiters();
}
}
}
#[cfg(test)]
mod tests {
use super::{TransportLifecycle, TransportState};
#[tokio::test]
async fn closing_waits_for_existing_admissions() {
let lifecycle = TransportLifecycle::new();
let admission = lifecycle.admit().unwrap();
lifecycle.begin_close();
assert!(lifecycle.admit().is_none());
drop(admission);
lifecycle.wait_admissions_drained().await;
lifecycle.mark_draining();
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Draining);
}
#[tokio::test]
async fn an_admitted_continuation_reopens_the_transport() {
let lifecycle = TransportLifecycle::new();
let admission = lifecycle.admit().unwrap();
lifecycle.begin_close();
drop(admission);
lifecycle.wait_admissions_drained().await;
lifecycle.mark_draining();
lifecycle.reopen();
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Open);
assert!(lifecycle.admit().is_some());
}
}
+2
View File
@@ -2,10 +2,12 @@
mod handle;
mod inbox;
mod lifecycle;
mod output;
mod registry;
pub use handle::*;
pub use inbox::*;
pub(crate) use lifecycle::*;
pub use output::*;
pub use registry::*;
+1 -13
View File
@@ -1,12 +1,11 @@
//! Buffers, replays, broadcasts, and atomically closes downstream output.
use bytes::Bytes;
use tokio::sync::{mpsc, Notify};
use tokio::sync::mpsc;
#[derive(Default)]
pub struct OutputHub {
state: parking_lot::Mutex<OutputState>,
closed: Notify,
}
#[derive(Default)]
@@ -49,17 +48,6 @@ impl OutputHub {
state.closed = true;
state.subscribers.clear();
drop(state);
self.closed.notify_waiters();
true
}
pub async fn wait_closed(&self) {
loop {
let notified = self.closed.notified();
if self.state.lock().closed {
return;
}
notified.await;
}
}
}
+76 -19
View File
@@ -1,13 +1,19 @@
//! Maps request IDs to active transport handles.
use std::{collections::HashMap, sync::Arc};
use std::{
collections::HashMap,
sync::{
atomic::{AtomicU64, Ordering},
Arc,
},
};
use tokio::sync::{mpsc, Mutex, Notify};
use crate::{
cursor::{
conversation::ConversationRegistry, prompting::PromptCompiler,
services::observability::CursorTraceRecorder,
services::observability::CursorTraceService,
},
plugin::PluginRegistry,
provider::Provider,
@@ -24,15 +30,23 @@ pub struct TransportRegistry {
}
struct RegistryInner {
local: Mutex<HashMap<String, TransportHandle>>,
local: Mutex<HashMap<String, LocalTransport>>,
next_local_generation: AtomicU64,
upstream: Mutex<HashMap<String, u64>>,
route_changed: Notify,
store: Store,
traces: CursorTraceService,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
conversations: ConversationRegistry,
}
#[derive(Clone)]
struct LocalTransport {
generation: u64,
handle: TransportHandle,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TransportRoute {
Local,
@@ -99,8 +113,10 @@ impl TransportRegistry {
Self {
inner: Arc::new(RegistryInner {
local: Mutex::new(HashMap::new()),
next_local_generation: AtomicU64::new(1),
upstream: Mutex::new(HashMap::new()),
route_changed: Notify::new(),
traces: CursorTraceService::new(store.clone()),
conversations: ConversationRegistry::new(
store.clone(),
provider,
@@ -119,6 +135,13 @@ impl TransportRegistry {
&self.inner.store
}
pub fn trace(
&self,
request_id: &str,
) -> crate::cursor::services::observability::CursorTraceRecorder {
self.inner.traces.recorder(request_id)
}
pub fn web_cache(&self) -> &WebCache {
&self.inner.web_cache
}
@@ -132,18 +155,37 @@ impl TransportRegistry {
}
pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> {
if let Some(handle) = self.inner.local.lock().await.get(request_id).cloned() {
return Ok(handle);
self.get_or_create_for_append(request_id, false).await
}
pub(crate) async fn get_or_create_for_append(
&self,
request_id: &str,
replace_closing: bool,
) -> Result<TransportHandle> {
let mut local = self.inner.local.lock().await;
if let Some(transport) = local.get(request_id) {
if transport.handle.accepting_appends() || !replace_closing {
return Ok(transport.handle.clone());
}
}
local.remove(request_id);
let (commands, receiver) = mpsc::channel(128);
let output = Arc::new(OutputHub::default());
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await;
let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace);
let mut local = self.inner.local.lock().await;
if let Some(existing) = local.get(request_id).cloned() {
return Ok(existing);
}
local.insert(request_id.into(), handle.clone());
let trace = self.inner.traces.recorder(request_id);
trace.resume();
let handle = TransportHandle::new(request_id.into(), commands, output, trace);
let generation = self
.inner
.next_local_generation
.fetch_add(1, Ordering::Relaxed);
local.insert(
request_id.into(),
LocalTransport {
generation,
handle: handle.clone(),
},
);
drop(local);
self.inner.route_changed.notify_waiters();
self.inner
@@ -152,17 +194,29 @@ impl TransportRegistry {
let registry = Arc::downgrade(&self.inner);
let request_id = request_id.to_string();
let lifecycle = handle.clone();
tokio::spawn(async move {
output.wait_closed().await;
lifecycle.wait_transport_closed().await;
if let Some(registry) = registry.upgrade() {
registry.local.lock().await.remove(&request_id);
let mut local = registry.local.lock().await;
if local
.get(&request_id)
.is_some_and(|transport| transport.generation == generation)
{
local.remove(&request_id);
}
}
});
Ok(handle)
}
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
self.inner.local.lock().await.get(request_id).cloned()
self.inner
.local
.lock()
.await
.get(request_id)
.map(|transport| transport.handle.clone())
}
pub async fn mark_upstream(&self, request_id: &str) {
@@ -206,10 +260,13 @@ impl TransportRegistry {
self.inner.conversations.shutdown().await;
let handles = std::mem::take(&mut *self.inner.local.lock().await);
self.inner.upstream.lock().await.clear();
for handle in handles.into_values() {
handle.disconnect().await;
let _ =
tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await;
for transport in handles.into_values() {
transport.handle.disconnect().await;
let _ = tokio::time::timeout(
std::time::Duration::from_secs(2),
transport.handle.wait_transport_closed(),
)
.await;
}
}
}
+75 -3
View File
@@ -52,23 +52,24 @@ async fn inject_if_missing_at(path: &Path) -> Result<()> {
.execute(&mut connection)
.await?;
let token = local_token()?;
let account = sqlx::query("SELECT CAST(value AS TEXT) AS value FROM ItemTable WHERE key = ?")
.bind("cursorAuth/accessToken")
.fetch_optional(&mut connection)
.await?;
if account.is_some_and(|row| {
row.try_get::<String, _>("value")
.is_ok_and(|value| !value.trim().is_empty())
.is_ok_and(|value| !value.trim().is_empty() && value != token)
}) {
return Ok(());
}
let token = local_token()?;
let values = [
("cursorAuth/accessToken", token.as_str()),
("cursorAuth/refreshToken", token.as_str()),
("cursorAuth/cachedEmail", EMAIL),
("cursorAuth/cachedSignUpType", SIGN_UP_TYPE),
("cursorAuth/stripeMembershipAuthId", SUBJECT),
("cursorAuth/stripeMembershipType", MEMBERSHIP_TYPE),
("cursorAuth/stripeSubscriptionStatus", SUBSCRIPTION_STATUS),
];
@@ -89,7 +90,17 @@ async fn inject_if_missing_at(path: &Path) -> Result<()> {
Ok(())
}
fn local_token() -> Result<String> {
pub(crate) fn is_local_cursor_authorization(authorization: &str) -> bool {
authorization
.strip_prefix("Bearer ")
.is_some_and(is_local_cursor_token)
}
fn is_local_cursor_token(token: &str) -> bool {
local_token().is_ok_and(|local| local == token)
}
pub(super) fn local_token() -> Result<String> {
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"HS256","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&json!({
"sub": SUBJECT,
@@ -101,3 +112,64 @@ fn local_token() -> Result<String> {
}))?);
Ok(format!("{header}.{payload}.{SUBJECT}"))
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn reinjection_repairs_local_membership_cache() {
let directory = tempfile::tempdir().unwrap();
let path = directory.path().join("state.vscdb");
inject_if_missing_at(&path).await.unwrap();
let mut connection = SqliteConnection::connect(&format!("sqlite:{}", path.display()))
.await
.unwrap();
sqlx::query("UPDATE ItemTable SET value = 'free' WHERE key = ?")
.bind("cursorAuth/stripeMembershipType")
.execute(&mut connection)
.await
.unwrap();
sqlx::query("DELETE FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipAuthId")
.execute(&mut connection)
.await
.unwrap();
drop(connection);
inject_if_missing_at(&path).await.unwrap();
let mut connection = SqliteConnection::connect(&format!("sqlite:{}", path.display()))
.await
.unwrap();
let membership_type: String =
sqlx::query_scalar("SELECT CAST(value AS TEXT) FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipType")
.fetch_one(&mut connection)
.await
.unwrap();
let membership_auth_id: String =
sqlx::query_scalar("SELECT CAST(value AS TEXT) FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipAuthId")
.fetch_one(&mut connection)
.await
.unwrap();
assert_eq!(membership_type, MEMBERSHIP_TYPE);
assert_eq!(membership_auth_id, SUBJECT);
}
#[test]
fn recognizes_only_the_injected_cursor_token() {
let token = local_token().unwrap();
assert!(is_local_cursor_token(&token));
assert!(is_local_cursor_authorization(&format!("Bearer {token}")));
assert!(!is_local_cursor_authorization(&token));
assert!(!is_local_cursor_authorization(
"Bearer official-cursor-token"
));
assert!(!is_local_cursor_token("official-cursor-token"));
assert!(!is_local_cursor_token(""));
}
}
+35 -3
View File
@@ -1,6 +1,7 @@
//! Exposes the local desktop application integration.
mod account;
mod ca;
mod process;
mod proxy;
mod settings;
@@ -21,8 +22,16 @@ pub(crate) fn proxy_host_allowed(host: &str) -> bool {
proxy::is_cursor_host(host)
}
fn integration_prerequisites_ready(ca: &CaState, backend_ready: bool) -> bool {
matches!(ca, CaState::Ready) && backend_ready
pub(crate) fn request_uses_local_cursor_token(headers: &axum::http::HeaderMap) -> bool {
headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.is_some_and(account::is_local_cursor_authorization)
}
#[cfg(test)]
pub(crate) fn local_cursor_authorization() -> String {
format!("Bearer {}", account::local_token().unwrap())
}
#[derive(Clone, Debug, Serialize)]
@@ -49,6 +58,7 @@ pub struct CursorHarnessStatus {
pub configured_models: usize,
pub enabled_models: usize,
pub integration: IntegrationState,
pub settings_applied: bool,
pub proxy_url: Option<String>,
pub ca_install_command: Option<String>,
}
@@ -90,6 +100,10 @@ impl CursorHarness {
*self.inner.backend_addr.write() = Some(addr);
}
pub async fn proxy_port(&self) -> Option<u16> {
self.inner.proxy.lock().await.port()
}
pub async fn cleanup_stale_settings(&self) -> Result<()> {
settings::clear_stale_managed_settings()
}
@@ -99,7 +113,10 @@ impl CursorHarness {
let configured_models = models.len();
let enabled_models = configured_models;
let ca = self.inner.ca.state()?;
if integration_prerequisites_ready(&ca, self.inner.backend_addr.read().is_some()) {
if self.inner.store.cursor_takeover_enabled().await?
&& matches!(ca, CaState::Ready)
&& self.inner.backend_addr.read().is_some()
{
self.enable().await?;
}
let proxy = self.inner.proxy.lock().await;
@@ -120,6 +137,7 @@ impl CursorHarness {
configured_models,
enabled_models,
integration,
settings_applied,
proxy_url,
ca_install_command: self.inner.ca.install_command(),
})
@@ -135,9 +153,23 @@ impl CursorHarness {
}
pub async fn set_enabled(&self, enabled: bool) -> Result<CursorHarnessStatus> {
let settings_applied = {
let proxy = self.inner.proxy.lock().await;
proxy
.url()
.as_deref()
.map(settings::settings_match)
.transpose()?
.unwrap_or(false)
};
if enabled {
if !settings_applied {
process::terminate_cursor().await?;
}
self.inner.store.set_cursor_takeover_enabled(true).await?;
self.enable().await?;
} else {
self.inner.store.set_cursor_takeover_enabled(false).await?;
self.disable().await?;
}
self.status().await
+79
View File
@@ -0,0 +1,79 @@
//! Terminates the Cursor desktop process before an explicit takeover.
use tokio::process::Command;
use crate::{Error, Result};
pub async fn terminate_cursor() -> Result<()> {
terminate_platform_cursor().await
}
#[cfg(target_os = "macos")]
async fn terminate_platform_cursor() -> Result<()> {
terminate_unix_process("Cursor").await
}
#[cfg(target_os = "linux")]
async fn terminate_platform_cursor() -> Result<()> {
terminate_unix_process("cursor").await?;
terminate_unix_process("Cursor").await
}
#[cfg(any(target_os = "macos", target_os = "linux"))]
async fn terminate_unix_process(name: &str) -> Result<()> {
let running = Command::new("pgrep").args(["-x", name]).status().await?;
if !running.success() {
return match running.code() {
Some(1) => Ok(()),
_ => Err(Error::Config(format!(
"failed to inspect the {name} process"
))),
};
}
let terminated = Command::new("pkill").args(["-x", name]).status().await?;
if terminated.success() || terminated.code() == Some(1) {
Ok(())
} else {
Err(Error::Config(format!(
"failed to terminate the {name} process"
)))
}
}
#[cfg(target_os = "windows")]
async fn terminate_platform_cursor() -> Result<()> {
let processes = Command::new("tasklist")
.args(["/FI", "IMAGENAME eq Cursor.exe", "/NH", "/FO", "CSV"])
.output()
.await?;
if !processes.status.success() {
return Err(Error::Config(
"failed to inspect the Cursor.exe process".into(),
));
}
if !String::from_utf8_lossy(&processes.stdout)
.to_ascii_lowercase()
.contains("cursor.exe")
{
return Ok(());
}
let terminated = Command::new("taskkill")
.args(["/F", "/T", "/IM", "Cursor.exe"])
.status()
.await?;
if terminated.success() {
Ok(())
} else {
Err(Error::Config(
"failed to terminate the Cursor.exe process".into(),
))
}
}
#[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))]
async fn terminate_platform_cursor() -> Result<()> {
Err(Error::Config(format!(
"terminating Cursor is unsupported on {}",
std::env::consts::OS
)))
}
+47
View File
@@ -33,6 +33,13 @@ impl ProxyRuntime {
pub fn url(&self) -> Option<String> {
self.running().then(|| self.url.clone()).flatten()
}
pub fn port(&self) -> Option<u16> {
if self.running() {
self.port
} else {
None
}
}
pub async fn start(
&mut self,
@@ -152,10 +159,21 @@ fn is_local_path(path: &str) -> bool {
path,
"/agent.v1.AgentService/RunSSE"
| "/aiserver.v1.BidiService/BidiAppend"
| "/aiserver.v1.AiService/AvailableDocs"
| "/aiserver.v1.DashboardService/GetEffectiveUserPlugins"
| "/aiserver.v1.DashboardService/GetUserPrivacyMode"
| "/agent.v1.AgentService/UpdateConversationMetadata"
| "/aiserver.v1.AiService/GetServerConfig"
| "/aiserver.v1.ServerConfigService/GetServerConfig"
| "/aiserver.v1.AiService/AvailableModels"
| "/agent.v1.AgentService/GetUsableModels"
| "/aiserver.v1.AiService/GetUsableModels"
| "/agent.v1.AgentService/GetDefaultModelForCli"
| "/aiserver.v1.AiService/GetDefaultModelForCli"
| "/aiserver.v1.AiService/GetDefaultModel"
| "/aiserver.v1.AiService/GetDefaultModelNudgeData"
| "/aiserver.v1.AuthService/GetEmail"
| "/aiserver.v1.AuthService/GetUserMeta"
| "/aiserver.v1.DashboardService/GetMe"
| "/aiserver.v1.DashboardService/GetTeams"
| "/aiserver.v1.DashboardService/GetUserProfile"
@@ -165,11 +183,40 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.AiService/KnowledgeBaseList"
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
| "/aiserver.v1.AiService/WriteGitCommitMessage"
| "/aiserver.v1.NetworkService/IsConnected"
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
| "/auth/full_stripe_profile"
| "/auth/stripe_profile"
)
}
fn should_route_locally(path: &str, tab_mode: TabMode) -> bool {
is_local_path(path) || (is_tab_path(path) && tab_mode != TabMode::Direct)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cursor_cli_transport_and_model_metadata_routes_stay_local() {
for path in [
"/aiserver.v1.AiService/GetServerConfig",
"/aiserver.v1.ServerConfigService/GetServerConfig",
"/agent.v1.AgentService/GetDefaultModelForCli",
"/aiserver.v1.AiService/GetDefaultModelForCli",
"/aiserver.v1.AiService/GetDefaultModel",
"/aiserver.v1.AiService/GetDefaultModelNudgeData",
"/aiserver.v1.AiService/AvailableDocs",
"/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
"/aiserver.v1.DashboardService/GetUserPrivacyMode",
"/aiserver.v1.AuthService/GetUserMeta",
"/agent.v1.AgentService/UpdateConversationMetadata",
"/auth/full_stripe_profile",
"/auth/stripe_profile",
] {
assert!(is_local_path(path), "{path} must not reach Cursor upstream");
}
}
}
+65 -1
View File
@@ -29,8 +29,72 @@ mod usage {
}
}
/// Adds two optional counts, treating an unreported count as zero.
/// The result is only unknown when neither side reported the field.
fn sum(left: Option<u64>, right: Option<u64>) -> Option<u64> {
left?.checked_add(right?)
if left.is_none() && right.is_none() {
return None;
}
Some(
left.unwrap_or_default()
.saturating_add(right.unwrap_or_default()),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_record_without_cache_details_keeps_the_counts_already_accumulated() {
let mut total = Usage {
input_tokens: Some(1_000),
output_tokens: Some(20),
total_tokens: Some(1_020),
cache_read_tokens: Some(900),
..Default::default()
};
// A second provider call that omits `prompt_tokens_details`.
total += Usage {
input_tokens: Some(1_200),
output_tokens: Some(30),
total_tokens: Some(1_230),
..Default::default()
};
assert_eq!(total.input_tokens, Some(2_200));
assert_eq!(total.output_tokens, Some(50));
assert_eq!(total.total_tokens, Some(2_250));
assert_eq!(total.cache_read_tokens, Some(900));
}
#[test]
fn a_late_reported_count_is_not_discarded() {
let mut total = Usage {
input_tokens: Some(1_000),
..Default::default()
};
total += Usage {
input_tokens: Some(1_200),
cache_read_tokens: Some(900),
reasoning_tokens: Some(64),
..Default::default()
};
assert_eq!(total.cache_read_tokens, Some(900));
assert_eq!(total.reasoning_tokens, Some(64));
}
#[test]
fn a_field_no_call_reported_stays_unknown() {
let mut total = Usage {
input_tokens: Some(1_000),
..Default::default()
};
total += Usage {
input_tokens: Some(1_200),
..Default::default()
};
assert_eq!(total.cache_write_tokens, None);
}
}
}
pub use usage::*;
+55 -6
View File
@@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize};
use crate::{Error, Result};
use super::{
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent,
ToolResultContent,
normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role,
ToolCallContent, ToolResultContent,
};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
@@ -91,7 +91,7 @@ fn project_tool_round(
"tool round repeats provider replay state".into(),
));
}
calls.extend(part_calls.iter().cloned());
calls.extend(part_calls.iter().map(normalized_tool_call));
cursor += 1;
while cursor < messages.len() {
@@ -139,7 +139,7 @@ fn project_tool_round(
.map(|(message_id, result)| ProjectedMessage {
message_id,
role: Role::Tool,
content: ProjectedContent::ToolResult(result),
content: ProjectedContent::ToolResult(normalized_tool_result(&result)),
}),
);
Ok(Some((output, cursor)))
@@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
text: text.clone(),
thinking: thinking.clone(),
replay_state: replay_state.clone(),
calls: tool_calls.clone(),
calls: tool_calls.iter().map(normalized_tool_call).collect(),
},
MessageContent::ToolResult(result) => ProjectedContent::ToolResult(result.clone()),
MessageContent::ToolResult(result) => {
ProjectedContent::ToolResult(normalized_tool_result(result))
}
};
ProjectedMessage {
message_id: message.message_id.clone(),
@@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
content,
}
}
fn normalized_tool_call(call: &ToolCallContent) -> ToolCallContent {
let mut call = call.clone();
call.name = normalize_tool_name(&call.name);
call
}
fn normalized_tool_result(result: &ToolResultContent) -> ToolResultContent {
let mut result = result.clone();
result.name = normalize_tool_name(&result.name);
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Origin;
use serde_json::json;
#[test]
fn tool_names_are_normalized_before_provider_dispatch() {
let messages = [CanonicalMessage {
message_id: "assistant-1".into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: String::new(),
thinking: String::new(),
tool_round_id: None,
replay_state: None,
tool_calls: vec![ToolCallContent {
index: 0,
call_id: "call-1".into(),
name: "multi_tool_use.parallel".into(),
arguments: json!({}),
}],
},
runtime_event_id: None,
}];
let projected = project_messages(&messages).unwrap();
let ProjectedContent::Assistant { calls, .. } = &projected[0].content else {
panic!("expected assistant projection");
};
assert_eq!(calls[0].name, "multi_tool_use_parallel");
}
}
+42 -11
View File
@@ -11,10 +11,15 @@ pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[Projected
let tools = prompt.tools.iter().fold(0_u64, |total, tool| {
total.saturating_add(estimate_json_tokens(tool))
});
let messages = messages.iter().fold(0_u64, |total, message| {
instructions
.saturating_add(tools)
.saturating_add(estimate_projected_messages_tokens(messages))
}
pub(crate) fn estimate_projected_messages_tokens(messages: &[ProjectedMessage]) -> u64 {
messages.iter().fold(0_u64, |total, message| {
total.saturating_add(estimate_message_tokens(message))
});
instructions.saturating_add(tools).saturating_add(messages)
})
}
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
@@ -23,7 +28,7 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
ProjectedContent::Assistant {
text,
thinking,
replay_state,
replay_state: _,
calls,
} => {
let calls = calls.iter().fold(0_u64, |total, call| {
@@ -35,12 +40,6 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
});
estimate_text_tokens(text)
.saturating_add(estimate_text_tokens(thinking))
.saturating_add(
replay_state
.as_ref()
.map(estimate_json_tokens)
.unwrap_or_default(),
)
.saturating_add(calls)
}
ProjectedContent::ToolResult(result) => {
@@ -118,7 +117,9 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent};
use crate::model::{
ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent,
};
fn prompt() -> PromptSpec {
PromptSpec {
@@ -214,4 +215,34 @@ mod tests {
assert!(with_text > with_call);
assert!(with_image > with_text);
}
#[test]
fn assistant_replay_state_does_not_duplicate_thinking_or_count_signature() {
let assistant = |replay_state| ProjectedMessage {
message_id: "assistant".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: "answer".into(),
thinking: "reasoning".repeat(1_000),
replay_state,
calls: Vec::new(),
},
};
let without_replay = assistant(None);
let with_replay = assistant(Some(ProviderReplayState {
provider_kind: "anthropic".into(),
value: serde_json::json!({
"blocks": [{
"type": "thinking",
"thinking": "reasoning".repeat(1_000),
"signature": "s".repeat(282_100)
}]
}),
}));
assert_eq!(
estimate_projected_messages_tokens(&[without_replay]),
estimate_projected_messages_tokens(&[with_replay])
);
}
}
+18
View File
@@ -4,6 +4,24 @@ use serde_json::Value;
use super::ProviderReplayState;
pub fn normalize_tool_name(name: &str) -> String {
let normalized = name
.chars()
.map(|character| {
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
character
} else {
'_'
}
})
.collect::<String>();
if normalized.is_empty() {
"_".into()
} else {
normalized
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
pub struct ToolDefinition {
pub name: String,
+202 -13
View File
@@ -1,7 +1,97 @@
//! Provides shared network client and transport configuration.
//! Outbound HTTP clients configured from persisted application proxy settings.
//! Owns reusable outbound HTTP clients configured from persisted proxy settings.
use crate::{store::Store, Result};
use std::{sync::Arc, time::Duration};
use tokio::sync::RwLock;
use crate::{
store::{ProxySettingsSecret, Store},
Error, Result,
};
const LOCAL_NO_PROXY: &str = "localhost,127.0.0.0/8,::1";
#[derive(Clone)]
pub struct NetworkClients {
store: Store,
cache: Arc<RwLock<ClientCache>>,
}
#[derive(Default)]
struct ClientCache {
default: Option<reqwest::Client>,
cursor: Option<reqwest::Client>,
provider: Option<(Duration, reqwest::Client)>,
}
impl NetworkClients {
pub fn new(store: Store) -> Self {
Self {
store,
cache: Arc::new(RwLock::new(ClientCache::default())),
}
}
pub async fn default_client(&self) -> Result<reqwest::Client> {
if let Some(client) = self.cache.read().await.default.clone() {
return Ok(client);
}
let mut cache = self.cache.write().await;
if let Some(client) = cache.default.clone() {
return Ok(client);
}
let client = client_builder(&self.store).await?.build()?;
cache.default = Some(client.clone());
Ok(client)
}
pub async fn cursor_client(&self) -> Result<reqwest::Client> {
if let Some(client) = self.cache.read().await.cursor.clone() {
return Ok(client);
}
let mut cache = self.cache.write().await;
if let Some(client) = cache.cursor.clone() {
return Ok(client);
}
let client = client_builder(&self.store)
.await?
.redirect(reqwest::redirect::Policy::none())
.build()?;
cache.cursor = Some(client.clone());
Ok(client)
}
pub async fn provider_client(&self, timeout: Duration) -> Result<reqwest::Client> {
if let Some((_, client)) = self
.cache
.read()
.await
.provider
.as_ref()
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
{
return Ok(client.clone());
}
let mut cache = self.cache.write().await;
if let Some((_, client)) = cache
.provider
.as_ref()
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
{
return Ok(client.clone());
}
let client = client_builder(&self.store)
.await?
.timeout(timeout)
.build()?;
cache.provider = Some((timeout, client.clone()));
Ok(client)
}
pub async fn invalidate(&self) {
*self.cache.write().await = ClientCache::default();
}
}
pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
let settings = store.proxy_settings_secret().await?;
@@ -9,11 +99,7 @@ pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
// only offer legacy TLS 1.2 cipher suites unsupported by rustls.
let mut builder = reqwest::Client::builder().use_native_tls();
if settings.mode.is_custom() {
let mut proxy = reqwest::Proxy::all(&settings.address)?;
if settings.auth_enabled {
proxy = proxy.basic_auth(&settings.username, &settings.password);
}
builder = builder.no_proxy().proxy(proxy);
builder = builder.proxy(custom_proxy(&settings)?);
}
Ok(builder)
}
@@ -26,11 +112,114 @@ pub async fn blocking_client_builder(store: &Store) -> Result<reqwest::blocking:
let settings = store.proxy_settings_secret().await?;
let mut builder = reqwest::blocking::Client::builder().use_native_tls();
if settings.mode.is_custom() {
let mut proxy = reqwest::Proxy::all(&settings.address)?;
if settings.auth_enabled {
proxy = proxy.basic_auth(&settings.username, &settings.password);
}
builder = builder.no_proxy().proxy(proxy);
builder = builder.proxy(custom_proxy(&settings)?);
}
Ok(builder)
}
fn custom_proxy(settings: &ProxySettingsSecret) -> Result<reqwest::Proxy> {
let mut proxy = reqwest::Proxy::all(&settings.address)?
.no_proxy(reqwest::NoProxy::from_string(LOCAL_NO_PROXY));
if settings.auth_enabled {
proxy = proxy.basic_auth(&settings.username, &settings.password);
}
Ok(proxy)
}
pub fn reject_self_proxy(address: &str, local_proxy_port: u16) -> Result<()> {
if local_proxy_port == 0 {
return Ok(());
}
let url = url::Url::parse(address)
.map_err(|error| Error::Config(format!("invalid proxy address: {error}")))?;
if url.port_or_known_default() == Some(local_proxy_port) && url_host_is_loopback(&url) {
return Err(Error::Config(
"proxy address cannot point to the Cursor BYOK local proxy".into(),
));
}
Ok(())
}
fn url_host_is_loopback(url: &url::Url) -> bool {
match url.host() {
Some(url::Host::Domain(host)) => {
host.trim_end_matches('.').eq_ignore_ascii_case("localhost")
}
Some(url::Host::Ipv4(address)) => address.is_loopback(),
Some(url::Host::Ipv6(address)) => address.is_loopback(),
None => false,
}
}
#[cfg(test)]
mod tests {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use super::*;
use crate::store::ProxyMode;
fn custom_settings(address: String) -> ProxySettingsSecret {
ProxySettingsSecret {
mode: ProxyMode::Custom,
address,
auth_enabled: false,
username: String::new(),
password: String::new(),
}
}
#[test]
fn rejects_only_own_loopback_proxy_port() {
for address in [
"http://localhost:15721",
"http://localhost.:15721",
"http://127.0.0.2:15721",
"http://[::1]:15721",
] {
assert!(reject_self_proxy(address, 15721).is_err(), "{address}");
}
assert!(reject_self_proxy("http://127.0.0.1:7890", 15721).is_ok());
assert!(reject_self_proxy("http://192.168.1.2:15721", 15721).is_ok());
assert!(reject_self_proxy("http://127.0.0.1:15721", 0).is_ok());
}
#[tokio::test]
async fn custom_proxy_bypasses_loopback_destinations() {
let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_address = proxy_listener.local_addr().unwrap();
let target_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let target_address = target_listener.local_addr().unwrap();
let target = tokio::spawn(async move {
let (mut socket, _) = target_listener.accept().await.unwrap();
let mut request = [0_u8; 1024];
let _ = socket.read(&mut request).await.unwrap();
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await
.unwrap();
});
let settings = custom_settings(format!("http://{proxy_address}"));
let client = reqwest::Client::builder()
.proxy(custom_proxy(&settings).unwrap())
.build()
.unwrap();
let body = client
.get(format!("http://{target_address}"))
.send()
.await
.unwrap()
.text()
.await
.unwrap();
assert_eq!(body, "ok");
target.await.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(50), proxy_listener.accept())
.await
.is_err(),
"loopback destination unexpectedly reached the configured proxy"
);
}
}
+12 -1
View File
@@ -145,6 +145,13 @@ const ANTIGRAVITY_AUTH: &[(&str, &str)] = &[
"/plugins/build-in/antigravity-auth/oauth.ts"
)),
),
(
"google_oauth.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/antigravity-auth/google_oauth.ts"
)),
),
(
"resources.ts",
include_str!(concat!(
@@ -162,9 +169,9 @@ const ANTIGRAVITY_AUTH: &[(&str, &str)] = &[
];
const PLUGINS: &[(&str, &[(&str, &str)])] = &[
("antigravity-auth", ANTIGRAVITY_AUTH),
("codex-auth", CODEX_AUTH),
("grok-auth", GROK_AUTH),
("antigravity-auth", ANTIGRAVITY_AUTH),
];
/// 把内置插件预装到 installed 目录。manifest 的 version 是缓存键:
@@ -265,6 +272,10 @@ mod tests {
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
embedded_main()
);
assert!(root
.path()
.join("antigravity-auth/assets/antigravity.svg")
.is_file());
// 版本一致:本地改动与额外文件保持原样,不发生任何写盘。
std::fs::write(plugin.join("main.ts"), "edited").unwrap();
+32 -1
View File
@@ -255,12 +255,43 @@ fn validate_definition(plugin_id: &str, definition: &PluginModuleDefinition) ->
}
for method in &resource.add {
validate_id(&method.id, "plugin add method id")?;
if method.method_type != super::descriptor::OAUTH2_ADD_METHOD {
if !matches!(
method.method_type.as_str(),
super::descriptor::OAUTH2_ADD_METHOD
| super::descriptor::OAUTH2_AUTHORIZATION_CODE_ADD_METHOD
) {
return Err(Error::Config(format!(
"plugin '{plugin_id}' add method '{}' uses unsupported type '{}'",
method.id, method.method_type
)));
}
if method.method_type == super::descriptor::OAUTH2_ADD_METHOD
&& method.callback.is_some()
{
return Err(Error::Config(format!(
"plugin '{plugin_id}' device OAuth method '{}' cannot declare callback settings",
method.id
)));
}
if let Some(callback) = &method.callback {
if callback.port == Some(0) {
return Err(Error::Config(format!(
"plugin '{plugin_id}' OAuth callback port must be greater than zero"
)));
}
if let Some(path) = callback.path.as_deref() {
if !path.starts_with('/')
|| path.len() > 128
|| path.contains('?')
|| path.contains('#')
|| path.contains("//")
{
return Err(Error::Config(format!(
"plugin '{plugin_id}' contains invalid OAuth callback path '{path}'"
)));
}
}
}
}
}
Ok(())
+12 -4
View File
@@ -203,20 +203,28 @@ fn validate_component(value: &str, label: &str) -> Result<()> {
Ok(())
}
fn set_directory_permissions(_path: &Path) -> Result<()> {
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))?;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700))?;
}
#[cfg(not(unix))]
{
let _ = path;
}
Ok(())
}
fn set_file_permissions(_path: &Path) -> Result<()> {
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))?;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
}
#[cfg(not(unix))]
{
let _ = path;
}
Ok(())
}
+12
View File
@@ -51,6 +51,17 @@ pub struct AddMethodDefinition {
pub display_name: LocalizedText,
#[serde(default)]
pub description: LocalizedText,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub callback: Option<OAuthCallbackDefinition>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct OAuthCallbackDefinition {
#[serde(default)]
pub port: Option<u16>,
#[serde(default)]
pub path: Option<String>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
@@ -64,6 +75,7 @@ pub struct ImportDefinition {
}
pub const OAUTH2_ADD_METHOD: &str = "oauth2.0";
pub const OAUTH2_AUTHORIZATION_CODE_ADD_METHOD: &str = "oauth2.authorization-code";
/// 桌面端看到的插件全貌。
#[derive(Clone, Debug, Serialize)]
+1 -2
View File
@@ -7,7 +7,7 @@ mod definition;
mod descriptor;
mod installation;
mod manifest;
pub mod oauth_callback;
mod oauth_callback;
mod protocol;
mod registry;
mod runtime;
@@ -21,7 +21,6 @@ pub use descriptor::{
};
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)]
+326 -77
View File
@@ -1,98 +1,347 @@
//! Lightweight local OAuth callback server for Google / Antigravity OAuth redirect flows.
use std::{collections::HashMap, net::SocketAddr, sync::Arc};
//! Core-owned loopback callback transport for plugin OAuth authorization-code flows.
use std::{net::SocketAddr, sync::Arc, time::Duration};
use axum::{
extract::{Query, State},
http::HeaderMap,
response::Html,
routing::get,
Json, Router,
Router,
};
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use serde::Deserialize;
use tokio::sync::{oneshot, Mutex};
use tokio_util::sync::CancellationToken;
#[derive(Default, Clone)]
pub struct OAuthCallbackState {
codes: Arc<RwLock<HashMap<String, String>>>,
use crate::{Error, Result};
const CALLBACK_RESPONSE_TIMEOUT: Duration = Duration::from_secs(120);
pub(super) struct CallbackRequest {
pub result: std::result::Result<String, String>,
pub response: oneshot::Sender<CallbackOutcome>,
}
#[derive(Deserialize)]
pub struct CallbackQuery {
pub code: Option<String>,
pub state: Option<String>,
pub error: Option<String>,
#[derive(Debug)]
pub(super) struct CallbackOutcome {
pub success: bool,
pub message: Option<String>,
}
#[derive(Deserialize)]
pub struct StatusQuery {
pub state: Option<String>,
pub(super) struct CallbackHandle {
pub redirect_uri: String,
pub receiver: oneshot::Receiver<CallbackRequest>,
shutdown: CancellationToken,
}
#[derive(Serialize)]
pub struct StatusResponse {
pub code: Option<String>,
}
pub async fn start_oauth_callback_server(shutdown: tokio_util::sync::CancellationToken) {
let port = 51121;
let addr = SocketAddr::from(([127, 0, 0, 1], port));
let state = OAuthCallbackState::default();
let router = Router::new()
.route("/oauth-callback", get(handle_callback))
.route("/auth-status", get(handle_status))
.with_state(state);
let listener = match tokio::net::TcpListener::bind(addr).await {
Ok(l) => l,
Err(err) => {
tracing::warn!(%addr, %err, "OAuth callback port 51121 unavailable or already bound");
return;
}
};
tracing::info!(%addr, "OAuth callback server listening");
let server = axum::serve(listener, router).with_graceful_shutdown(async move {
shutdown.cancelled().await;
});
if let Err(err) = server.await {
tracing::debug!(%err, "OAuth callback server stopped");
impl Drop for CallbackHandle {
fn drop(&mut self) {
self.shutdown.cancel();
}
}
#[derive(Clone)]
struct CallbackState {
expected_state: String,
plugin_name: String,
plugin_icon: String,
resource_name: serde_json::Value,
sender: Arc<Mutex<Option<oneshot::Sender<CallbackRequest>>>>,
shutdown: CancellationToken,
}
#[derive(Deserialize)]
struct CallbackQuery {
code: Option<String>,
state: Option<String>,
error: Option<String>,
error_description: Option<String>,
}
pub(super) async fn bind(
port: Option<u16>,
path: &str,
expected_state: String,
plugin_name: String,
plugin_icon: String,
resource_name: serde_json::Value,
) -> Result<CallbackHandle> {
let address = SocketAddr::from(([127, 0, 0, 1], port.unwrap_or(0)));
let listener = tokio::net::TcpListener::bind(address)
.await
.map_err(|error| {
Error::Config(format!(
"cannot bind plugin OAuth callback at {address}: {error}"
))
})?;
let local_address = listener.local_addr()?;
let redirect_uri = format!("http://127.0.0.1:{}{path}", local_address.port());
let (sender, receiver) = oneshot::channel();
let shutdown = CancellationToken::new();
let state = CallbackState {
expected_state,
plugin_name,
plugin_icon,
resource_name,
sender: Arc::new(Mutex::new(Some(sender))),
shutdown: shutdown.clone(),
};
let router = Router::new()
.route(path, get(handle_callback))
.with_state(state);
let graceful = shutdown.clone();
tokio::spawn(async move {
if let Err(error) = axum::serve(listener, router)
.with_graceful_shutdown(async move { graceful.cancelled().await })
.await
{
tracing::debug!(%error, "plugin OAuth callback server stopped");
}
});
Ok(CallbackHandle {
redirect_uri,
receiver,
shutdown,
})
}
async fn handle_callback(
State(state): State<OAuthCallbackState>,
State(state): State<CallbackState>,
headers: HeaderMap,
Query(query): Query<CallbackQuery>,
) -> Html<&'static str> {
if let (Some(code), Some(st)) = (query.code, query.state) {
state.codes.write().insert(st, code);
) -> Html<String> {
let locale = callback_locale(&headers);
if query.state.as_deref() != Some(state.expected_state.as_str()) {
return Html(render_page(
&state,
locale,
false,
Some(localized(
locale,
"授权状态不匹配,请返回应用后重试。",
"Authorization state did not match. Return to the app and try again.",
)),
));
}
let result = match query.code.filter(|code| !code.trim().is_empty()) {
Some(code) => Ok(code),
None => Err(query.error_description.or(query.error).unwrap_or_else(|| {
localized(locale, "授权被取消。", "Authorization was cancelled.").to_owned()
})),
};
let Some(sender) = state.sender.lock().await.take() else {
return Html(render_page(
&state,
locale,
false,
Some(localized(
locale,
"该授权回调已被使用。",
"This authorization callback has already been used.",
)),
));
};
let (response, completion) = oneshot::channel();
if sender.send(CallbackRequest { result, response }).is_err() {
return Html(render_page(
&state,
locale,
false,
Some(localized(
locale,
"授权会话已结束。",
"The authorization session has ended.",
)),
));
}
let outcome = tokio::time::timeout(CALLBACK_RESPONSE_TIMEOUT, completion).await;
state.shutdown.cancel();
match outcome {
Ok(Ok(outcome)) => Html(render_page(
&state,
locale,
outcome.success,
outcome.message.as_deref(),
)),
_ => Html(render_page(
&state,
locale,
false,
Some(localized(
locale,
"添加资源超时,请返回应用后重试。",
"Adding the resource timed out. Return to the app and try again.",
)),
)),
}
Html(r#"<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
<title>Antigravity Authorization Successful</title>
<style>
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif; background: #18181b; color: #f4f4f5; display: flex; align-items: center; justify-content: center; height: 100vh; margin: 0; }
.card { background: #27272a; border: 1px solid #3f3f46; padding: 36px; border-radius: 16px; text-align: center; box-shadow: 0 10px 30px rgba(0,0,0,0.5); max-width: 440px; }
.icon { font-size: 40px; color: #4285f4; margin-bottom: 16px; }
h1 { color: #ffffff; font-size: 20px; margin: 0 0 12px 0; }
p { color: #a1a1aa; font-size: 14px; line-height: 1.6; margin: 0; }
</style>
</head>
<body>
<div class="card">
<div class="icon">✓</div>
<h1>Authorization Successful</h1>
<p>Your Google Antigravity account has been authorized. You can close this browser tab and return to <strong>Cursor BYOK</strong>.</p>
</div>
</body>
</html>"#)
}
async fn handle_status(
State(state): State<OAuthCallbackState>,
Query(query): Query<StatusQuery>,
) -> Json<StatusResponse> {
let code = query.state.and_then(|st| state.codes.read().get(&st).cloned());
Json(StatusResponse { code })
fn callback_locale(headers: &HeaderMap) -> &'static str {
headers
.get(axum::http::header::ACCEPT_LANGUAGE)
.and_then(|value| value.to_str().ok())
.filter(|value| value.to_ascii_lowercase().starts_with("zh"))
.map(|_| "zh-CN")
.unwrap_or("en-US")
}
fn localized<'a>(locale: &str, chinese: &'a str, english: &'a str) -> &'a str {
if locale == "zh-CN" {
chinese
} else {
english
}
}
fn localized_value(value: &serde_json::Value, locale: &str) -> String {
value
.as_str()
.map(str::to_owned)
.or_else(|| {
value
.get(locale)
.and_then(serde_json::Value::as_str)
.map(str::to_owned)
})
.or_else(|| {
value
.get("en-US")
.and_then(serde_json::Value::as_str)
.map(str::to_owned)
})
.or_else(|| {
value
.as_object()?
.values()
.find_map(serde_json::Value::as_str)
.map(str::to_owned)
})
.unwrap_or_else(|| localized(locale, "资源", "resource").to_owned())
}
fn escape_html(value: &str) -> String {
value
.replace('&', "&amp;")
.replace('<', "&lt;")
.replace('>', "&gt;")
.replace('"', "&quot;")
.replace('\'', "&#39;")
}
fn render_page(
state: &CallbackState,
locale: &str,
success: bool,
message: Option<&str>,
) -> String {
let plugin_name = escape_html(&state.plugin_name);
let plugin_icon = escape_html(&state.plugin_icon);
let resource_name = escape_html(&localized_value(&state.resource_name, locale));
let title = if success {
localized(locale, "资源添加成功", "Resource added")
} else {
localized(locale, "资源添加失败", "Could not add resource")
};
let detail = message.map(escape_html).unwrap_or_else(|| {
if success {
localized(
locale,
"已为该插件添加资源。",
"A resource has been added for this plugin.",
)
.to_owned()
} else {
localized(
locale,
"请返回应用后重试。",
"Return to the app and try again.",
)
.to_owned()
}
});
let close = localized(
locale,
"您现在可以关闭本页面并返回 Cursor BYOK。",
"You can now close this page and return to Cursor BYOK.",
);
format!(
r#"<!doctype html>
<html lang="{locale}"><head><meta charset="utf-8"><meta name="viewport" content="width=device-width,initial-scale=1">
<meta http-equiv="Content-Security-Policy" content="default-src 'none'; img-src data:; style-src 'unsafe-inline'">
<title>{title}</title><style>
:root{{color-scheme:light dark}}body{{margin:0;min-height:100vh;display:grid;place-items:center;font:15px system-ui,-apple-system,sans-serif;background:Canvas;color:CanvasText}}main{{width:min(420px,calc(100vw - 48px));text-align:center}}img{{width:56px;height:56px;object-fit:contain}}h1{{font-size:20px;margin:16px 0 6px}}.plugin{{opacity:.72;margin-bottom:24px}}.resource{{font-weight:600;margin:8px 0}}.detail{{opacity:.82;line-height:1.6}}.close{{opacity:.62;margin-top:24px;font-size:13px}}
</style></head><body><main><img src="{plugin_icon}" alt=""><h1>{title}</h1><div class="plugin">{plugin_name}</div><div class="resource">{resource_name}</div><div class="detail">{detail}</div><div class="close">{close}</div></main></body></html>"#
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn escapes_plugin_content_in_callback_page() {
let state = CallbackState {
expected_state: "state".into(),
plugin_name: "<plugin>".into(),
plugin_icon: "data:image/svg+xml;base64,abc".into(),
resource_name: serde_json::json!({"en-US": "Accounts & keys"}),
sender: Arc::new(Mutex::new(None)),
shutdown: CancellationToken::new(),
};
let page = render_page(&state, "en-US", true, None);
assert!(page.contains("&lt;plugin&gt;"));
assert!(page.contains("Accounts &amp; keys"));
assert!(!page.contains("<plugin>"));
}
#[tokio::test]
async fn callback_rejects_wrong_state_then_delivers_code() {
let mut callback = bind(
None,
"/oauth-callback",
"expected".into(),
"Plugin".into(),
"data:image/svg+xml;base64,abc".into(),
serde_json::json!("Account"),
)
.await
.unwrap();
let client = reqwest::Client::new();
let rejected = client
.get(format!(
"{}?state=wrong&code=ignored",
callback.redirect_uri
))
.send()
.await
.unwrap()
.text()
.await
.unwrap();
assert!(rejected.contains("Authorization state did not match"));
let client = client.clone();
let redirect_uri = callback.redirect_uri.clone();
let browser = tokio::spawn(async move {
client
.get(format!("{redirect_uri}?state=expected&code=accepted"))
.send()
.await
.unwrap()
.text()
.await
.unwrap()
});
let request = (&mut callback.receiver).await.unwrap();
assert_eq!(request.result.unwrap(), "accepted");
request
.response
.send(CallbackOutcome {
success: true,
message: None,
})
.unwrap();
assert!(browser.await.unwrap().contains("Resource added"));
}
}
+396 -128
View File
@@ -1,8 +1,10 @@
//! Orchestrates plugin capabilities: resources, model catalogs, and invocation.
use std::{collections::HashMap, path::Path, sync::Arc};
use std::{collections::HashMap, path::Path, sync::Arc, time::Duration};
use async_stream::try_stream;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
use serde::Serialize;
use sha2::{Digest, Sha256};
use tokio::sync::{Mutex, RwLock};
use tokio_util::sync::CancellationToken;
@@ -12,16 +14,20 @@ use super::{
descriptor::{
parse_model_id, PluginDescriptor, PluginModelDescriptor, PluginProviderDescriptor,
PluginResourceDescriptor, PluginResourceView, ProviderDefinition, ResourceDefinition,
ResourcePresentation, OAUTH2_ADD_METHOD,
ResourcePresentation, OAUTH2_ADD_METHOD, OAUTH2_AUTHORIZATION_CODE_ADD_METHOD,
},
oauth_callback::{self, CallbackHandle, CallbackOutcome, CallbackRequest},
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,
model::ModelInvocation,
provider::ProviderStream,
provider::{CallRecorder, ModelEvent},
store::Store,
Error, Result,
};
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
@@ -40,7 +46,6 @@ struct RegistryInner {
entries: RwLock<Option<Vec<PluginEntry>>>,
workers: Mutex<HashMap<String, Arc<PluginWorker>>>,
oauth_sessions: Mutex<HashMap<String, OAuthSession>>,
rr_counter: std::sync::atomic::AtomicUsize,
}
struct OAuthSession {
@@ -51,13 +56,42 @@ struct OAuthSession {
expires_at_ms: i64,
poll_interval_ms: i64,
next_poll_at_ms: i64,
flow: OAuthFlow,
}
enum OAuthFlow {
DeviceCode,
AuthorizationCode {
redirect_uri: String,
code_verifier: String,
callback: CallbackHandle,
},
}
enum OAuthPollWork {
DeviceCode {
plugin_id: String,
resource_type: String,
method_id: String,
session: serde_json::Value,
poll_interval_ms: i64,
},
AuthorizationCode {
plugin_id: String,
resource_type: String,
method_id: String,
session: serde_json::Value,
redirect_uri: String,
code_verifier: String,
callback_request: CallbackRequest,
},
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct OAuthBeginResponse {
pub session_id: String,
pub user_code: String,
pub user_code: Option<String>,
pub verification_url: String,
pub verification_url_complete: Option<String>,
pub expires_at_ms: i64,
@@ -109,7 +143,6 @@ impl PluginRegistry {
entries: RwLock::new(None),
workers: Mutex::new(HashMap::new()),
oauth_sessions: Mutex::new(HashMap::new()),
rr_counter: std::sync::atomic::AtomicUsize::new(0),
}),
})
}
@@ -144,12 +177,6 @@ impl PluginRegistry {
let Some(executable) = self.inner.runtime.executable() else {
return Vec::new();
};
let disabled_models = self
.inner
.store
.disabled_plugin_models()
.await
.unwrap_or_default();
let mut models = Vec::new();
for entry in self.entries(&executable).await {
for provider in &entry.definition.providers {
@@ -162,19 +189,14 @@ impl PluginRegistry {
.models(&entry.manifest.id, &provider.id)
.await
.unwrap_or_default();
models.extend(stored.iter().filter_map(|model| {
let descriptor = PluginModelDescriptor::new(
models.extend(stored.iter().map(|model| {
PluginModelDescriptor::new(
&entry.manifest.id,
&entry.manifest.name,
&entry.icon,
provider,
model,
);
if disabled_models.contains(&descriptor.id) {
None
} else {
Some(descriptor)
}
)
}));
}
}
@@ -205,15 +227,6 @@ impl PluginRegistry {
}
pub async fn plan_model(&self, model_id: &str) -> Result<PluginInvocationPlan> {
let disabled_models = self
.inner
.store
.disabled_plugin_models()
.await
.unwrap_or_default();
if disabled_models.contains(model_id) {
return Err(Error::Provider(format!("plugin model '{model_id}' is disabled")));
}
let model = self.model_descriptor(model_id).await?;
let request_url = format!("plugin://{}/{}", model.plugin_id, model.provider_id);
Ok(PluginInvocationPlan { model, request_url })
@@ -225,6 +238,7 @@ impl PluginRegistry {
&self,
invocation: ModelInvocation,
cancellation: CancellationToken,
recorder: CallRecorder,
) -> ProviderStream {
let registry = self.clone();
Box::pin(try_stream! {
@@ -254,7 +268,7 @@ impl PluginRegistry {
"request": request,
});
let worker = registry.worker(&entry, &executable).await;
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone(), Some(recorder)).await?;
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
while let Some(item) = items.recv().await {
match item {
@@ -306,48 +320,145 @@ impl PluginRegistry {
let method = resource
.add
.iter()
.find(|method| method.id == method_id && method.method_type == OAUTH2_ADD_METHOD)
.find(|method| method.id == method_id)
.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 worker = self.worker(&entry, &executable).await;
let session_id = uuid::Uuid::new_v4().to_string();
let (
session,
verification_url,
verification_url_complete,
expires_at_ms,
poll_interval_ms,
flow,
user_code,
) = match method.method_type.as_str() {
OAUTH2_ADD_METHOD => {
let value = worker
.invoke(
"oauth.begin",
serde_json::json!({
"resourceType": resource_type,
"methodId": method.id,
}),
CancellationToken::new(),
)
.await?;
let begin: OAuth2Begin = serde_json::from_value(value)?;
(
begin.session,
begin.verification_url,
begin.verification_url_complete,
begin.expires_at_ms,
begin.poll_interval_ms.max(1_000),
OAuthFlow::DeviceCode,
Some(begin.user_code),
)
}
OAUTH2_AUTHORIZATION_CODE_ADD_METHOD => {
let state = oauth_random_secret();
let code_verifier = oauth_random_secret();
let code_challenge =
URL_SAFE_NO_PAD.encode(Sha256::digest(code_verifier.as_bytes()));
let callback = method.callback.as_ref();
let callback = oauth_callback::bind(
callback.and_then(|value| value.port),
callback
.and_then(|value| value.path.as_deref())
.unwrap_or("/oauth-callback"),
state.clone(),
entry.manifest.name.clone(),
entry.icon.clone(),
serde_json::to_value(&resource.display_name)?,
)
.await?;
let redirect_uri = callback.redirect_uri.clone();
let value = worker
.invoke(
"oauth.begin",
serde_json::json!({
"resourceType": resource_type,
"methodId": method.id,
"authorization": {
"redirectUri": redirect_uri,
"state": state,
"codeChallenge": code_challenge,
},
}),
CancellationToken::new(),
)
.await?;
let begin: OAuth2AuthorizationCodeBegin = serde_json::from_value(value)?;
let poll_interval_ms = begin.poll_interval_ms.unwrap_or(1_000).max(1_000);
(
begin.session,
begin.authorization_url,
None,
begin.expires_at_ms,
poll_interval_ms,
OAuthFlow::AuthorizationCode {
redirect_uri,
code_verifier,
callback,
},
None,
)
}
method_type => {
return Err(Error::Config(format!(
"plugin '{plugin_id}' OAuth method '{method_id}' uses unsupported type '{method_type}'"
)));
}
};
if expires_at_ms <= now_ms() {
return Err(Error::Protocol(format!(
"plugin '{plugin_id}' OAuth method '{method_id}' returned an expired session"
)));
}
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),
session,
expires_at_ms,
poll_interval_ms,
next_poll_at_ms: now_ms() + poll_interval_ms,
flow,
},
);
let cleanup = self.clone();
let cleanup_session_id = session_id.clone();
let cleanup_delay_ms = expires_at_ms.saturating_sub(now_ms()) as u64;
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(cleanup_delay_ms)).await;
cleanup
.inner
.oauth_sessions
.lock()
.await
.remove(&cleanup_session_id);
});
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),
user_code,
verification_url,
verification_url_complete,
expires_at_ms,
poll_interval_ms,
})
}
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 work = {
let now = now_ms();
let mut sessions = self.inner.oauth_sessions.lock().await;
let Some(state) = sessions.get_mut(session_id) else {
return Ok(OAuthPollResponse::Failed {
@@ -357,7 +468,7 @@ impl PluginRegistry {
if now >= state.expires_at_ms {
sessions.remove(session_id);
return Ok(OAuthPollResponse::Failed {
message: "device authorization expired".into(),
message: "authorization expired".into(),
});
}
if now < state.next_poll_at_ms {
@@ -366,14 +477,101 @@ impl PluginRegistry {
});
}
state.next_poll_at_ms = now + state.poll_interval_ms;
(
let common = (
state.plugin_id.clone(),
state.resource_type.clone(),
state.method_id.clone(),
state.session.clone(),
state.poll_interval_ms,
)
);
match &mut state.flow {
OAuthFlow::DeviceCode => OAuthPollWork::DeviceCode {
plugin_id: common.0,
resource_type: common.1,
method_id: common.2,
session: common.3,
poll_interval_ms: state.poll_interval_ms,
},
OAuthFlow::AuthorizationCode {
redirect_uri,
code_verifier,
callback,
} => match callback.receiver.try_recv() {
Ok(callback_request) => OAuthPollWork::AuthorizationCode {
plugin_id: common.0,
resource_type: common.1,
method_id: common.2,
session: common.3,
redirect_uri: redirect_uri.clone(),
code_verifier: code_verifier.clone(),
callback_request,
},
Err(tokio::sync::oneshot::error::TryRecvError::Empty) => {
return Ok(OAuthPollResponse::Pending {
poll_interval_ms: state.poll_interval_ms,
});
}
Err(tokio::sync::oneshot::error::TryRecvError::Closed) => {
sessions.remove(session_id);
return Ok(OAuthPollResponse::Failed {
message: "authorization callback stopped before completion".into(),
});
}
},
}
};
match work {
OAuthPollWork::DeviceCode {
plugin_id,
resource_type,
method_id,
session,
poll_interval_ms,
} => {
self.poll_device_code(
session_id,
plugin_id,
resource_type,
method_id,
session,
poll_interval_ms,
)
.await
}
OAuthPollWork::AuthorizationCode {
plugin_id,
resource_type,
method_id,
session,
redirect_uri,
code_verifier,
callback_request,
} => {
self.complete_authorization_code(
session_id,
plugin_id,
resource_type,
method_id,
session,
redirect_uri,
code_verifier,
callback_request,
)
.await
}
}
}
async fn poll_device_code(
&self,
session_id: &str,
plugin_id: String,
resource_type: String,
method_id: String,
session: serde_json::Value,
poll_interval_ms: i64,
) -> Result<OAuthPollResponse> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, &plugin_id).await?;
let value = self
@@ -404,21 +602,18 @@ impl PluginRegistry {
})
}
OAuth2Poll::Completed { resources } => {
// 持久化成功后才销毁会话:写盘瞬时失败时下次轮询还能重试。
let outcome = self
.inner
.state
.upsert_resources(&plugin_id, &resource_type, resources)
// 持久化成功后才销毁设备码会话:写盘瞬时失败时下次轮询还能重试。
let response = self
.persist_oauth_resources(
&entry,
&executable,
&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,
})
Ok(response)
}
OAuth2Poll::Denied { message } => {
self.inner.oauth_sessions.lock().await.remove(session_id);
@@ -431,6 +626,103 @@ impl PluginRegistry {
}
}
#[allow(clippy::too_many_arguments)]
async fn complete_authorization_code(
&self,
session_id: &str,
plugin_id: String,
resource_type: String,
method_id: String,
session: serde_json::Value,
redirect_uri: String,
code_verifier: String,
callback_request: CallbackRequest,
) -> Result<OAuthPollResponse> {
let CallbackRequest { result, response } = callback_request;
let code = match result {
Ok(code) => code,
Err(message) => {
self.inner.oauth_sessions.lock().await.remove(session_id);
let _ = response.send(CallbackOutcome {
success: false,
message: Some(message.clone()),
});
return Ok(OAuthPollResponse::Denied {
message: Some(message),
});
}
};
let result = async {
let executable = self.executable()?;
let entry = self.find_entry(&executable, &plugin_id).await?;
let value = self
.worker(&entry, &executable)
.await
.invoke(
"oauth.complete",
serde_json::json!({
"resourceType": resource_type,
"methodId": method_id,
"session": session,
"authorization": {
"code": code,
"redirectUri": redirect_uri,
"codeVerifier": code_verifier,
},
}),
CancellationToken::new(),
)
.await?;
let resources: Vec<ResourceDraft> = serde_json::from_value(value)?;
self.persist_oauth_resources(&entry, &executable, &plugin_id, &resource_type, resources)
.await
}
.await;
self.inner.oauth_sessions.lock().await.remove(session_id);
match result {
Ok(completed) => {
let _ = response.send(CallbackOutcome {
success: true,
message: None,
});
Ok(completed)
}
Err(error) => {
let message = error.to_string();
let _ = response.send(CallbackOutcome {
success: false,
message: Some(message.clone()),
});
Ok(OAuthPollResponse::Failed { message })
}
}
}
async fn persist_oauth_resources(
&self,
entry: &PluginEntry,
executable: &Path,
plugin_id: &str,
resource_type: &str,
resources: Vec<ResourceDraft>,
) -> Result<OAuthPollResponse> {
let outcome = self
.inner
.state
.upsert_resources(plugin_id, resource_type, resources)
.await?;
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,
})
}
pub async fn import_resources(
&self,
plugin_id: &str,
@@ -726,21 +1018,13 @@ impl PluginRegistry {
}
}
match &provider.resource_type {
Some(resource_type) => {
let disabled_accounts = self
.inner
.store
.disabled_plugin_accounts()
.await
.unwrap_or_default();
let resources = self
.inner
.state
.resources(plugin_id, resource_type)
.await
.unwrap_or_default();
resources.iter().any(|r| !disabled_accounts.contains(&r.id))
}
Some(resource_type) => !self
.inner
.state
.resources(plugin_id, resource_type)
.await
.unwrap_or_default()
.is_empty(),
None => true,
}
}
@@ -826,54 +1110,19 @@ impl PluginRegistry {
plugin_id: &str,
resource_type: &str,
) -> Result<ResourceRecord> {
let disabled_accounts = self
.inner
.store
.disabled_plugin_accounts()
.await
.unwrap_or_default();
let records = self.inner.state.resources(plugin_id, resource_type).await?;
let active_records: Vec<_> = records
.into_iter()
.filter(|record| !disabled_accounts.contains(&record.id))
.collect();
if active_records.is_empty() {
if records.is_empty() {
return Err(Error::Provider(format!(
"plugin '{plugin_id}' has no enabled '{resource_type}' resource; enable or add one first"
"plugin '{plugin_id}' has no '{resource_type}' resource; add one first"
)));
}
let now = now_ms();
let mut ready_records: Vec<_> = active_records
.into_iter()
.filter(|record| record.state.is_ready(now))
.collect();
if ready_records.is_empty() {
return Err(Error::Provider(format!(
"all enabled accounts for plugin '{plugin_id}' are currently cooling or rate-limited"
)));
}
let get_priority = |r: &ResourceRecord| -> u8 {
let label = r.private_data.get("quota").and_then(|q| q.get("planLabel")).and_then(|l| l.as_str()).unwrap_or("");
let lower = label.to_lowercase();
if label.contains("🔥") || lower.contains("pro") || lower.contains("ultra") || lower.contains("premium") || lower.contains("advanced") {
0
} else {
1
}
};
ready_records.sort_by_key(|r| get_priority(r));
if let Some(best_prio) = ready_records.first().map(|r| get_priority(r)) {
ready_records.retain(|r| get_priority(r) == best_prio);
}
let index = self
.inner
.rr_counter
.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
% ready_records.len();
Ok(ready_records[index].clone())
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(
@@ -960,6 +1209,25 @@ struct OAuth2Begin {
poll_interval_ms: i64,
}
#[derive(Debug, serde::Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct OAuth2AuthorizationCodeBegin {
session: serde_json::Value,
authorization_url: String,
expires_at_ms: i64,
#[serde(default)]
poll_interval_ms: Option<i64>,
}
fn oauth_random_secret() -> String {
let first = uuid::Uuid::new_v4();
let second = uuid::Uuid::new_v4();
let mut bytes = [0_u8; 32];
bytes[..16].copy_from_slice(first.as_bytes());
bytes[16..].copy_from_slice(second.as_bytes());
URL_SAFE_NO_PAD.encode(bytes)
}
#[derive(Debug, serde::Deserialize)]
#[serde(rename_all = "kebab-case", tag = "status")]
enum OAuth2Poll {
+6
View File
@@ -84,6 +84,12 @@ export function __descriptor(definition: ProviderPluginDefinition) {
id: method.id,
displayName: method.displayName,
description: method.description ?? null,
callback: method.type === "oauth2.authorization-code"
? {
port: method.callback?.port ?? null,
path: method.callback?.path ?? "/oauth-callback",
}
: null,
})),
import: resource.import
? {
@@ -102,6 +102,7 @@ export function buildChatBody(call: OpenAiChatCall): Record<string, JsonValue> {
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 };
}
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
if (!Array.isArray(items)) {
throw new Error("OpenAI Responses replay state is missing items");
}
return items;
return items.map((item) => {
const source = record(item);
if (source?.type !== "reasoning") {
throw new Error("OpenAI Responses replay state contains a non-reasoning item");
}
const projected: Record<string, JsonValue> = { type: "reasoning" };
for (const field of ["id", "summary", "content", "encrypted_content"] as const) {
if (field in source) projected[field] = source[field] as JsonValue;
}
return projected;
});
}
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
+36 -2
View File
@@ -66,7 +66,7 @@ export type OAuth2AddMethod = {
};
export type OAuth2Begin = {
/** 不透明流程状态(设备码、PKCE verifier 等);永远不会持久化。 */
/** 不透明流程状态(如设备码);永远不会持久化。 */
session: JsonValue;
userCode: string;
verificationUrl: string;
@@ -82,7 +82,41 @@ export type OAuth2Poll =
| { status: "denied"; message?: string }
| { status: "failed"; message: string };
export type ResourceAddMethod = OAuth2AddMethod;
/** Core 托管浏览器回调、state 与 PKCE 的 OAuth 2.0 授权码流程。 */
export type OAuth2AuthorizationCodeAddMethod = {
type: "oauth2.authorization-code";
id: string;
displayName: LocalizedText;
description?: LocalizedText;
/** 仅在上游 OAuth 客户端要求固定 loopback 地址时指定。 */
callback?: { port?: number; path?: string };
begin(
input: {
redirectUri: string;
state: string;
codeChallenge: string;
},
context: PluginContext,
): Promise<OAuth2AuthorizationCodeBegin>;
complete(
session: JsonValue,
input: {
code: string;
redirectUri: string;
codeVerifier: string;
},
context: PluginContext,
): Promise<ResourceDraft[]>;
};
export type OAuth2AuthorizationCodeBegin = {
session: JsonValue;
authorizationUrl: string;
expiresAtMs: number;
pollIntervalMs?: number;
};
export type ResourceAddMethod = OAuth2AddMethod | OAuth2AuthorizationCodeAddMethod;
export type ResourceImportFile = {
name: string;
+20 -2
View File
@@ -130,12 +130,30 @@ async function dispatch(message: { id: string; method: string; params?: JsonValu
}
case "oauth.begin": {
const support = resourceSupport(params.resourceType);
result = await addMethod(support, params.methodId).begin(context);
const method = addMethod(support, params.methodId);
result = method.type === "oauth2.authorization-code"
? await method.begin(params.authorization as never, context)
: await method.begin(context);
break;
}
case "oauth.poll": {
const support = resourceSupport(params.resourceType);
result = await addMethod(support, params.methodId).poll(params.session ?? null, context);
const method = addMethod(support, params.methodId);
if (method.type !== "oauth2.0") throw new Error(`add method ${method.id} does not support polling`);
result = await method.poll(params.session ?? null, context);
break;
}
case "oauth.complete": {
const support = resourceSupport(params.resourceType);
const method = addMethod(support, params.methodId);
if (method.type !== "oauth2.authorization-code") {
throw new Error(`add method ${method.id} does not support authorization-code completion`);
}
result = await method.complete(
params.session ?? null,
params.authorization as never,
context,
);
break;
}
case "import.parse": {
+271 -33
View File
@@ -3,7 +3,10 @@ use std::{
collections::{HashMap, HashSet},
path::PathBuf,
process::Stdio,
sync::Arc,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::Duration,
};
@@ -19,7 +22,7 @@ use super::{
definition::{file_url, PluginDefinitionLoader},
protocol::{HostMessage, WorkerMessage},
};
use crate::{store::Store, Error, Result};
use crate::{provider::CallRecorder, store::Store, Error, Result};
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
@@ -56,12 +59,29 @@ struct WorkerProcess {
stdin: Arc<Mutex<ChildStdin>>,
}
struct InvocationState {
cancellation: CancellationToken,
recorder: Option<CallRecorder>,
recorder_claimed: AtomicBool,
}
impl InvocationState {
fn claim_recorder(&self) -> Option<CallRecorder> {
self.recorder.as_ref().and_then(|recorder| {
self.recorder_claimed
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.ok()
.map(|_| recorder.clone())
})
}
}
#[derive(Clone)]
struct HostContext {
plugin_id: String,
network_hosts: Arc<HashSet<String>>,
store: Store,
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
}
@@ -87,7 +107,7 @@ impl PluginWorker {
.collect(),
),
store,
cancellations: Arc::new(Mutex::new(HashMap::new())),
invocations: Arc::new(Mutex::new(HashMap::new())),
streams: Arc::new(Mutex::new(HashMap::new())),
},
plugin_id,
@@ -108,7 +128,9 @@ impl PluginWorker {
params: serde_json::Value,
cancellation: CancellationToken,
) -> Result<serde_json::Value> {
let mut items = self.invoke_streaming(method, params, cancellation).await?;
let mut items = self
.invoke_streaming(method, params, cancellation, None)
.await?;
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
while let Some(item) = items.recv().await {
if let WorkerStreamItem::Result(result) = item {
@@ -137,15 +159,18 @@ impl PluginWorker {
method: &str,
params: serde_json::Value,
cancellation: CancellationToken,
recorder: Option<CallRecorder>,
) -> 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());
self.inner.host.invocations.lock().await.insert(
id.clone(),
Arc::new(InvocationState {
cancellation: request_cancellation.clone(),
recorder,
recorder_claimed: AtomicBool::new(false),
}),
);
let (sender, receiver) = mpsc::unbounded_channel();
self.inner
.pending
@@ -182,10 +207,10 @@ impl PluginWorker {
}
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
inner.pending.lock().await.remove(&request_id);
inner.host.cancellations.lock().await.remove(&request_id);
inner.host.invocations.lock().await.remove(&request_id);
}
_ = sender.closed() => {
inner.host.cancellations.lock().await.remove(&request_id);
inner.host.invocations.lock().await.remove(&request_id);
}
}
});
@@ -201,7 +226,7 @@ impl PluginWorker {
async fn cleanup(&self, id: &str) {
self.inner.pending.lock().await.remove(id);
self.inner.host.cancellations.lock().await.remove(id);
self.inner.host.invocations.lock().await.remove(id);
}
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
@@ -393,6 +418,28 @@ async fn fail_pending(pending: &Pending, message: &str) {
}
}
fn recorded_network_request(
params: &serde_json::Value,
) -> Result<(serde_json::Value, serde_json::Value)> {
let mut recorded_headers = serde_json::Map::new();
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"))
})?;
if !crate::model::is_sensitive_header(name) {
recorded_headers.insert(name.clone(), value.into());
}
}
}
let body = params
.get("body")
.and_then(serde_json::Value::as_str)
.map(|body| serde_json::from_str(body).unwrap_or_else(|_| body.into()))
.unwrap_or(serde_json::Value::Null);
Ok((serde_json::Value::Object(recorded_headers), body))
}
impl HostContext {
async fn call(
&self,
@@ -421,20 +468,17 @@ impl HostContext {
&self,
request_id: &str,
params: &serde_json::Value,
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
) -> Result<(
reqwest::RequestBuilder,
CancellationToken,
Option<CallRecorder>,
)> {
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}")))?;
let is_loopback = url
.host_str()
.map(|h| h == "127.0.0.1" || h == "localhost")
.unwrap_or(false);
if (url.scheme() != "https" && (!is_loopback || url.scheme() != "http"))
|| !url.username().is_empty()
|| url.password().is_some()
{
if url.scheme() != "https" || !url.username().is_empty() || url.password().is_some() {
return Err(Error::Config(
"plugin network URL must be HTTPS without credentials (or loopback HTTP)".into(),
"plugin network URL must be HTTPS without credentials".into(),
));
}
let host = url
@@ -470,14 +514,17 @@ impl HostContext {
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()
let invocation = self.invocations.lock().await.get(request_id).cloned();
let cancellation = invocation
.as_ref()
.map(|state| state.cancellation.clone())
.unwrap_or_default();
Ok((request, cancellation))
let recorder = invocation.and_then(|state| state.claim_recorder());
if let Some(recorder) = &recorder {
let (headers, body) = recorded_network_request(params)?;
recorder.request(headers, &body).await?;
}
Ok((request, cancellation, recorder))
}
async fn fetch(
@@ -485,13 +532,16 @@ impl HostContext {
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let (request, cancellation, recorder) = 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 let Some(recorder) = &recorder {
recorder.response_headers(status).await?;
}
if response
.content_length()
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
@@ -510,6 +560,9 @@ impl HostContext {
"plugin network response is larger than allowed".into(),
));
}
if let Some(recorder) = &recorder {
recorder.response_chunk(&body).await?;
}
Ok(
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
)
@@ -521,12 +574,15 @@ impl HostContext {
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let (request, cancellation, recorder) = 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();
if let Some(recorder) = &recorder {
recorder.response_headers(status).await?;
}
let headers = header_map(&response);
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
tokio::spawn(async move {
@@ -559,6 +615,12 @@ impl HostContext {
.await;
return;
}
if let Some(recorder) = &recorder {
if let Err(error) = recorder.response_chunk(&chunk).await {
let _ = sender.send(Err(error)).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>>();
@@ -652,3 +714,179 @@ fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a s
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
}
#[cfg(test)]
mod tests {
use crate::{
model::{NewLlmCall, ProviderType},
provider::{CallRecorder, FinishReason},
store::Store,
};
use super::*;
async fn recorder(detailed: bool, call_id: &str) -> (tempfile::TempDir, Store, CallRecorder) {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
store.set_detailed_logging(detailed).await.unwrap();
let recorder = CallRecorder::start(
store.clone(),
NewLlmCall {
call_id: call_id.into(),
run_id: "run".into(),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: "plugin:test/provider/model".into(),
provider_type: ProviderType::Plugin,
provider_url: "plugin://test/provider".into(),
request_type: ProviderType::Plugin,
request_url: "plugin://test/provider".into(),
model_id: "model".into(),
display_name: "Model".into(),
reasoning_effort: None,
fast: false,
message_count: 1,
tool_count: 0,
detailed: false,
},
)
.await
.unwrap();
(directory, store, recorder)
}
fn network_params() -> serde_json::Value {
serde_json::json!({
"url": "https://example.com/v1/responses",
"method": "POST",
"headers": {
"Authorization": "Bearer secret",
"X-Api-Key": "secret-key",
"Cookie": "session=secret",
"content-type": "application/json",
"x-client-request-id": "request-1"
},
"body": "{\"model\":\"test\",\"stream\":true}"
})
}
async fn host_with_recorder(store: Store, recorder: CallRecorder) -> HostContext {
let invocations = Arc::new(Mutex::new(HashMap::new()));
invocations.lock().await.insert(
"invocation".into(),
Arc::new(InvocationState {
cancellation: CancellationToken::new(),
recorder: Some(recorder),
recorder_claimed: AtomicBool::new(false),
}),
);
HostContext {
plugin_id: "test".into(),
network_hosts: Arc::new(HashSet::from(["example.com".into()])),
store,
invocations,
streams: Arc::new(Mutex::new(HashMap::new())),
}
}
#[test]
fn recorded_plugin_request_omits_sensitive_headers() {
let (headers, body) = recorded_network_request(&network_params()).unwrap();
assert_eq!(
headers,
serde_json::json!({
"content-type": "application/json",
"x-client-request-id": "request-1"
})
);
assert_eq!(body, serde_json::json!({ "model": "test", "stream": true }));
}
#[tokio::test]
async fn detailed_plugin_network_recording_persists_request_and_raw_response() {
let (_directory, store, recorder) = recorder(true, "detailed-plugin").await;
let host = host_with_recorder(store.clone(), recorder.clone()).await;
let params = network_params();
let (_, _, first_recorder) = host.request("invocation", &params).await.unwrap();
let (_, _, second_recorder) = host.request("invocation", &params).await.unwrap();
let (_, body) = recorded_network_request(&params).unwrap();
assert!(first_recorder.is_some());
assert!(second_recorder.is_none());
recorder.response_headers(200).await.unwrap();
recorder
.response_chunk(b"data: {\"type\":\"response.created\"}\n\n")
.await
.unwrap();
recorder.response_chunk(b"data: [DONE]\n\n").await.unwrap();
recorder.completed(FinishReason::Stop).await.unwrap();
let request = store
.llm_call_request("detailed-plugin")
.await
.unwrap()
.unwrap();
assert_eq!(
request.headers,
serde_json::json!({
"content-type": "application/json",
"x-client-request-id": "request-1"
})
);
assert_eq!(request.body, body);
let chunks = store.llm_call_chunks("detailed-plugin").await.unwrap();
let expected_response = "data: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n";
assert_eq!(chunks.len(), 2);
assert_eq!(
chunks
.iter()
.map(|chunk| chunk.data.as_str())
.collect::<String>(),
expected_response
);
let summary = store.llm_call("detailed-plugin").await.unwrap().unwrap();
assert_eq!(summary.http_status, Some(200));
assert_eq!(summary.stream_event_count, 2);
assert_eq!(summary.response_bytes, expected_response.len() as i64);
assert!(summary.detailed);
}
#[tokio::test]
async fn standard_plugin_network_recording_keeps_metrics_without_payloads() {
let (_directory, store, recorder) = recorder(false, "standard-plugin").await;
let host = host_with_recorder(store.clone(), recorder.clone()).await;
let params = network_params();
let (_, body) = recorded_network_request(&params).unwrap();
let request_bytes = serde_json::to_string(&body).unwrap().len() as i64;
let response = b"data: [DONE]\n\n";
let (_, _, observed) = host.request("invocation", &params).await.unwrap();
assert!(observed.is_some());
recorder.response_headers(204).await.unwrap();
recorder.response_chunk(response).await.unwrap();
recorder.completed(FinishReason::Stop).await.unwrap();
assert!(store
.llm_call_request("standard-plugin")
.await
.unwrap()
.is_none());
assert!(store
.llm_call_chunks("standard-plugin")
.await
.unwrap()
.is_empty());
let summary = store.llm_call("standard-plugin").await.unwrap().unwrap();
assert_eq!(summary.http_status, Some(204));
assert_eq!(summary.request_bytes, Some(request_bytes));
assert_eq!(summary.response_bytes, response.len() as i64);
assert_eq!(summary.stream_event_count, 1);
assert!(!summary.detailed);
}
}
+1
View File
@@ -6,6 +6,7 @@ mod normalize;
mod openai_chat;
mod openai_responses;
mod recorder;
mod request_template;
mod router;
use std::pin::Pin;
+35 -1
View File
@@ -365,7 +365,9 @@ fn map_finish(value: &str, has_tools: bool) -> FinishReason {
match value {
"tool_calls" | "function_call" => FinishReason::ToolUse,
"length" => FinishReason::Length,
"stop" | "content_filter" => FinishReason::Stop,
"content_filter" => FinishReason::Stop,
// Observed tool calls outrank ordinary or unknown stop labels because
// OpenAI-compatible servers commonly emit tools with `stop`.
_ if has_tools => FinishReason::ToolUse,
_ => FinishReason::Stop,
}
@@ -433,3 +435,35 @@ pub(crate) fn openai_usage(value: &Value) -> Usage {
.and_then(Value::as_u64),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn observed_tool_calls_outrank_ordinary_stop_reasons() {
assert_eq!(map_finish("stop", true), FinishReason::ToolUse);
assert_eq!(map_finish("", true), FinishReason::ToolUse);
}
#[test]
fn content_filter_never_executes_observed_tool_calls() {
assert_eq!(map_finish("content_filter", true), FinishReason::Stop);
assert_eq!(map_finish("content_filter", false), FinishReason::Stop);
}
#[test]
fn finish_reason_mapping_without_tool_calls_is_unchanged() {
assert_eq!(map_finish("stop", false), FinishReason::Stop);
assert_eq!(map_finish("content_filter", false), FinishReason::Stop);
assert_eq!(map_finish("", false), FinishReason::Stop);
assert_eq!(map_finish("tool_calls", false), FinishReason::ToolUse);
assert_eq!(map_finish("function_call", false), FinishReason::ToolUse);
}
#[test]
fn a_truncated_response_stays_truncated_even_with_tool_calls() {
assert_eq!(map_finish("length", true), FinishReason::Length);
assert_eq!(map_finish("length", false), FinishReason::Length);
}
}
+69 -4
View File
@@ -151,9 +151,8 @@ impl Provider for OpenAiResponsesProvider {
if !thinking_open { thinking_open = true; yield ModelEvent::ThinkingStart; }
if let Some(delta) = value.get("delta").and_then(Value::as_str) { yield ModelEvent::ThinkingDelta(delta.into()); }
}
"response.reasoning_summary_text.done" | "response.reasoning_text.done" => {
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
}
"response.reasoning_summary_text.done" | "response.reasoning_text.done"
if thinking_open => { thinking_open = false; yield ModelEvent::ThinkingEnd; }
"response.output_item.added" => {
let item = value.get("item").unwrap_or(&Value::Null);
if item.get("type").and_then(Value::as_str) == Some("function_call") {
@@ -420,7 +419,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
.ok_or_else(|| {
Error::Protocol("OpenAI Responses replay state is missing items".into())
})?;
input.extend(items.iter().cloned());
input.extend(
items
.iter()
.map(response_reasoning_input)
.collect::<Result<Vec<_>>>()?,
);
}
push_responses_text(&mut input, &message.role, text);
for call in calls {
@@ -437,6 +441,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
Ok(input)
}
fn response_reasoning_input(item: &Value) -> Result<Value> {
let source = item
.as_object()
.filter(|object| object.get("type").and_then(Value::as_str) == Some("reasoning"))
.ok_or_else(|| {
Error::Protocol("OpenAI Responses replay state contains a non-reasoning item".into())
})?;
let mut projected = Map::new();
projected.insert("type".into(), json!("reasoning"));
for field in ["id", "summary", "content", "encrypted_content"] {
if let Some(value) = source.get(field) {
projected.insert(field.into(), value.clone());
}
}
Ok(Value::Object(projected))
}
fn push_responses_parts(input: &mut Vec<Value>, role: &Role, parts: &[ContentPart]) -> Result<()> {
let text_type = if *role == Role::Assistant {
"output_text"
@@ -517,3 +538,47 @@ fn responses_usage(value: &Value) -> Usage {
.and_then(Value::as_u64),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ProviderReplayState;
#[test]
fn reasoning_replay_projects_response_items_to_valid_input_items() {
let messages = [ProjectedMessage {
message_id: "assistant-1".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: String::new(),
thinking: String::new(),
replay_state: Some(ProviderReplayState {
provider_kind: "openai_responses".into(),
value: json!({
"items": [{
"type": "reasoning",
"id": "item-1",
"status": "completed",
"summary": [{"type": "summary_text", "text": "why"}],
"content": [],
"encrypted_content": "opaque",
"output_only": true
}]
}),
}),
calls: Vec::new(),
},
}];
assert_eq!(
responses_input(&messages).unwrap(),
vec![json!({
"type": "reasoning",
"id": "item-1",
"summary": [{"type": "summary_text", "text": "why"}],
"content": [],
"encrypted_content": "opaque"
})]
);
}
}
+67
View File
@@ -0,0 +1,67 @@
//! Renders per-invocation placeholders in user-configured provider request values.
const SESSION_ID_PLACEHOLDER: &str = "{{SessionId}}";
pub(super) fn render_json_strings(
value: &serde_json::Value,
conversation_id: &str,
) -> serde_json::Value {
match value {
serde_json::Value::String(value) => {
serde_json::Value::String(value.replace(SESSION_ID_PLACEHOLDER, conversation_id))
}
serde_json::Value::Array(values) => serde_json::Value::Array(
values
.iter()
.map(|value| render_json_strings(value, conversation_id))
.collect(),
),
serde_json::Value::Object(values) => serde_json::Value::Object(
values
.iter()
.map(|(name, value)| (name.clone(), render_json_strings(value, conversation_id)))
.collect(),
),
value => value.clone(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn renders_session_id_in_nested_json_string_values() {
let template = serde_json::json!({
"session": "{{SessionId}}",
"nested": {
"label": "conversation={{SessionId}}/{{SessionId}}",
"values": ["{{SessionId}}", 42, true, null]
}
});
assert_eq!(
render_json_strings(&template, "cursor-conversation-id"),
serde_json::json!({
"session": "cursor-conversation-id",
"nested": {
"label": "conversation=cursor-conversation-id/cursor-conversation-id",
"values": ["cursor-conversation-id", 42, true, null]
}
})
);
}
#[test]
fn leaves_json_property_names_and_unrelated_strings_unchanged() {
let template = serde_json::json!({
"{{SessionId}}": "literal",
"other": "{{sessionId}}"
});
assert_eq!(
render_json_strings(&template, "cursor-conversation-id"),
template
);
}
}
+48 -7
View File
@@ -21,6 +21,7 @@ use super::{
pub struct ProviderRouter {
store: Store,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
request_timeout: Duration,
stream_idle_timeout: Duration,
}
@@ -29,12 +30,14 @@ impl ProviderRouter {
pub fn new(
store: Store,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
request_timeout: Duration,
stream_idle_timeout: Duration,
) -> Self {
Self {
store,
plugins,
clients,
request_timeout,
stream_idle_timeout,
}
@@ -49,6 +52,7 @@ impl Provider for ProviderRouter {
) -> ProviderStream {
let store = self.store.clone();
let plugins = self.plugins.clone();
let clients = self.clients.clone();
let request_timeout = self.request_timeout;
let stream_idle_timeout = self.stream_idle_timeout;
Box::pin(try_stream! {
@@ -62,7 +66,6 @@ impl Provider for ProviderRouter {
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 {
@@ -70,6 +73,7 @@ impl Provider for ProviderRouter {
}
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
registry: plugins.clone(),
recorder: recorder.clone(),
})));
(recorder, guard, provider.stream(routed, cancellation.clone()))
} else {
@@ -78,7 +82,11 @@ impl Provider for ProviderRouter {
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.extra_params =
super::request_template::render_json_strings(
model.extra_params(),
&invocation.conversation_id,
);
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();
@@ -86,12 +94,19 @@ impl Provider for ProviderRouter {
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() },
custom_headers: if model.custom_headers_enabled {
custom_headers(
&model.custom_headers,
&invocation.conversation_id,
)?
} else {
reqwest::header::HeaderMap::new()
},
max_output_tokens: model.max_output_tokens(),
request_timeout,
allowed_body_fields: None,
};
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
let client = clients.provider_client(request_timeout).await?;
let provider = build_observed(&config, recorder.clone(), client)?;
(recorder, guard, provider.stream(routed, cancellation.clone()))
};
@@ -227,6 +242,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
struct PluginModelProvider {
registry: PluginRegistry,
recorder: CallRecorder,
}
impl Provider for PluginModelProvider {
@@ -235,7 +251,8 @@ impl Provider for PluginModelProvider {
invocation: ModelInvocation,
cancellation: CancellationToken,
) -> ProviderStream {
self.registry.stream_model(invocation, cancellation)
self.registry
.stream_model(invocation, cancellation, self.recorder.clone())
}
}
@@ -291,8 +308,12 @@ fn root_error_message(error: &(dyn std::error::Error + 'static)) -> String {
current.to_string()
}
fn custom_headers(value: &serde_json::Value) -> Result<reqwest::header::HeaderMap> {
let object = value
fn custom_headers(
value: &serde_json::Value,
conversation_id: &str,
) -> Result<reqwest::header::HeaderMap> {
let rendered = super::request_template::render_json_strings(value, conversation_id);
let object = rendered
.as_object()
.ok_or_else(|| Error::Config("custom headers must be an object".into()))?;
let mut headers = reqwest::header::HeaderMap::new();
@@ -350,6 +371,26 @@ fn build_inner(
mod tests {
use super::*;
#[test]
fn renders_cursor_conversation_id_in_custom_header_values() {
let template = serde_json::json!({
"x-opencode-session-id": "{{SessionId}}",
"x-label": "cursor/{{SessionId}}"
});
let headers = custom_headers(&template, "cursor-conversation-id").unwrap();
assert_eq!(
headers.get("x-opencode-session-id").unwrap(),
"cursor-conversation-id"
);
assert_eq!(
headers.get("x-label").unwrap(),
"cursor/cursor-conversation-id"
);
assert_eq!(template["x-opencode-session-id"], "{{SessionId}}");
}
#[tokio::test]
async fn pending_provider_event_hits_the_idle_timeout() {
let mut stream: ProviderStream = Box::pin(futures_util::stream::pending());
+118 -10
View File
@@ -2,7 +2,13 @@
use std::collections::HashSet;
use crate::model::{estimate_context_tokens, CanonicalMessage, PreparedRun, ProjectedMessage};
use crate::{
model::{
estimate_context_tokens, estimate_projected_messages_tokens, CanonicalMessage, PreparedRun,
ProjectedMessage,
},
store::ContextUsageAnchor,
};
const FALLBACK_CHARS: usize = 12_000;
@@ -20,25 +26,44 @@ pub(super) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
pub(super) fn estimated_tokens(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> u64 {
estimate_context_tokens(&prepared.prompt, projected_messages)
anchor
.filter(|anchor| anchor.message_count <= projected_messages.len())
.map(|anchor| {
anchor
.context_input_tokens
.saturating_add(estimate_projected_messages_tokens(
&projected_messages[anchor.message_count..],
))
})
.unwrap_or_else(|| estimate_context_tokens(&prepared.prompt, projected_messages))
}
pub(super) fn compaction_estimate(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> Option<u64> {
let budget = input_budget(prepared)?;
let estimated = estimated_tokens(prepared, projected_messages, anchor);
(estimated > budget).then_some(estimated)
}
#[cfg(test)]
pub(super) fn should_compact(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> bool {
let Some(budget) = input_budget(prepared) else {
return false;
};
estimated_tokens(prepared, projected_messages) > budget
compaction_estimate(prepared, projected_messages, anchor).is_some()
}
pub(super) fn validate_compacted(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
) -> std::result::Result<u64, String> {
let estimated = estimated_tokens(prepared, projected_messages);
let estimated = estimate_context_tokens(&prepared.prompt, projected_messages);
let Some(budget) = input_budget(prepared) else {
return Ok(estimated);
};
@@ -125,14 +150,97 @@ mod tests {
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
let mut prepared = prepared(estimated + RESERVE_TOKENS);
assert!(!should_compact(&prepared, &projected));
assert!(!should_compact(&prepared, &projected, None));
prepared.model.context_window_tokens = Some(estimated + RESERVE_TOKENS - 1);
assert!(should_compact(&prepared, &projected));
assert!(should_compact(&prepared, &projected, None));
prepared.action = RunAction::Resume {
pending_tool_round: None,
};
assert!(should_compact(&prepared, &projected));
assert!(should_compact(&prepared, &projected, None));
}
#[test]
fn provider_usage_anchor_only_estimates_messages_added_after_last_request() {
let messages = vec![
CanonicalMessage::text("old", Role::User, Origin::Runtime, "x".repeat(400_000)),
CanonicalMessage::text("new", Role::User, Origin::Runtime, "short follow-up"),
];
let projected = project_messages(&messages).unwrap();
let anchor = ContextUsageAnchor {
context_input_tokens: 103_904,
message_count: 1,
};
let expected = 103_904 + estimate_projected_messages_tokens(&projected[1..]);
assert_eq!(
estimated_tokens(&prepared(200_000), &projected, Some(anchor)),
expected
);
assert!(!should_compact(
&prepared(200_000),
&projected,
Some(anchor)
));
}
#[test]
fn provider_usage_anchor_triggers_after_new_messages_cross_budget() {
let messages = vec![
CanonicalMessage::text("old", Role::User, Origin::Runtime, "old"),
CanonicalMessage::text("new", Role::User, Origin::Runtime, "x".repeat(80_000)),
];
let projected = project_messages(&messages).unwrap();
assert!(should_compact(
&prepared(200_000),
&projected,
Some(ContextUsageAnchor {
context_input_tokens: 180_000,
message_count: 1,
})
));
}
#[test]
fn missing_anchor_uses_full_fallback() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let prepared = prepared(200_000);
assert_eq!(
estimated_tokens(&prepared, &projected, None),
estimate_context_tokens(&prepared.prompt, &projected)
);
}
#[test]
fn invalid_anchor_message_count_uses_full_fallback() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let expected = estimate_context_tokens(&prepared(200_000).prompt, &projected);
assert_eq!(
estimated_tokens(
&prepared(200_000),
&projected,
Some(ContextUsageAnchor {
context_input_tokens: 1,
message_count: 2,
})
),
expected
);
}
#[test]
+73 -4
View File
@@ -10,7 +10,7 @@ use crate::{
ToolRoundId, Usage,
},
provider::Provider,
store::{RunStatus, Store},
store::{ContextUsageAnchor, RunStatus, Store},
};
use super::{
@@ -92,6 +92,14 @@ impl RunEngine {
cancellation: &CancellationToken,
) -> (RunOutcome, Option<Usage>) {
let mut usage = None;
let mut context_usage_anchor = match self
.store
.latest_context_usage(prepared.conversation_id.as_str())
.await
{
Ok(anchor) => anchor,
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
tracing::info!(
checkpoint_id = checkpoint.0,
"Run claimed conversation ownership"
@@ -175,15 +183,28 @@ impl RunEngine {
Ok(history) => history,
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
if prepared.action != RunAction::Compact
&& super::compaction::should_compact(prepared, &history)
{
let compaction_estimate = (prepared.action != RunAction::Compact)
.then(|| {
super::compaction::compaction_estimate(prepared, &history, context_usage_anchor)
})
.flatten();
if let Some(estimated_tokens) = compaction_estimate {
if emit(
client,
RunEvent::UsageSnapshot(context_usage_snapshot(estimated_tokens)),
)
.await
.is_err()
{
return (client_failure(), usage);
}
match self
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
.await
{
Ok((next_checkpoint, compaction_usage)) => {
checkpoint = next_checkpoint;
context_usage_anchor = None;
if let Some(compaction_usage) = compaction_usage {
accumulate_usage(&mut usage, compaction_usage);
}
@@ -272,11 +293,21 @@ impl RunEngine {
match interrupted {
Ok(cycle) => {
if let Some(cycle_usage) = cycle.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
}
Err(failure) => {
if let Some(cycle_usage) = failure.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
}
@@ -319,6 +350,11 @@ impl RunEngine {
Ok(cycle) => break 'attempt cycle,
Err(cycle_failure) => {
if let Some(cycle_usage) = cycle_failure.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
if cancellation.is_cancelled() {
@@ -420,6 +456,11 @@ impl RunEngine {
}
};
if let Some(cycle_usage) = cycle.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
@@ -819,6 +860,9 @@ impl RunEngine {
emit(client, RunEvent::AutoCompactionCompleted)
.await
.map_err(|_| client_failure())?;
emit(client, RunEvent::UsageSnapshot(context_usage_snapshot(0)))
.await
.map_err(|_| client_failure())?;
checkpoint = super::messages::append_batches(
&self.store,
prepared,
@@ -879,6 +923,31 @@ async fn hydrate_tool_images(
Ok(())
}
fn context_usage_snapshot(tokens: u64) -> Usage {
Usage {
input_tokens: Some(tokens),
context_input_tokens: Some(tokens),
output_tokens: Some(0),
total_tokens: Some(tokens),
cache_read_tokens: Some(0),
cache_write_tokens: Some(0),
reasoning_tokens: Some(0),
}
}
fn update_context_usage_anchor(
anchor: &mut Option<ContextUsageAnchor>,
usage: Usage,
message_count: usize,
) {
if let Some(context_input_tokens) = usage.context_input_tokens {
*anchor = Some(ContextUsageAnchor {
context_input_tokens,
message_count,
});
}
}
fn accumulate_usage(total: &mut Option<Usage>, usage: Usage) {
match total {
Some(total) => *total += usage,
+1
View File
@@ -129,6 +129,7 @@ pub enum RunEvent {
ToolCallEnd {
index: usize,
},
UsageSnapshot(Usage),
Usage(Usage),
ExecuteToolRound {
round_id: ToolRoundId,
+37 -8
View File
@@ -6,7 +6,7 @@ use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::{
model::{ProviderReplayState, ToolCall, Usage},
model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage},
provider::{FinishReason, ModelEvent, ProviderStream},
};
@@ -41,7 +41,7 @@ pub async fn consume_model_cycle(
mut stream: ProviderStream,
client: &mpsc::Sender<RunEvent>,
cancellation: &CancellationToken,
) -> std::result::Result<ModelCycleResult, ModelCycleFailure> {
) -> std::result::Result<ModelCycleResult, Box<ModelCycleFailure>> {
let mut model_call_id = None;
let mut text = String::new();
let mut reasoning = String::new();
@@ -165,6 +165,7 @@ pub async fn consume_model_cycle(
call_id,
name,
} => {
let name = normalize_tool_name(&name);
let Some(model_call_id) = model_call_id.as_ref() else {
return Err(failure(
RunFailure::Protocol("provider emitted content before Start".into()),
@@ -371,15 +372,15 @@ fn failure(
partial_text: String,
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
) -> Box<ModelCycleFailure> {
let retryable = matches!(failure, RunFailure::Protocol(_) | RunFailure::Provider(_));
ModelCycleFailure {
Box::new(ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable,
}
})
}
fn terminal_failure(
@@ -387,14 +388,14 @@ fn terminal_failure(
partial_text: String,
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
ModelCycleFailure {
) -> Box<ModelCycleFailure> {
Box::new(ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable: false,
}
})
}
#[cfg(test)]
@@ -406,6 +407,34 @@ mod tests {
};
use tokio_stream::wrappers::ReceiverStream;
#[tokio::test]
async fn provider_tool_names_are_normalized_when_received() {
let events = vec![
Ok(ModelEvent::Start {
model_call_id: "call".into(),
}),
Ok(ModelEvent::ToolCallStart {
index: 0,
call_id: "tool-call".into(),
name: "multi_tool_use.parallel".into(),
}),
Ok(ModelEvent::ToolCallEnd { index: 0 }),
Ok(ModelEvent::Done(FinishReason::ToolUse)),
];
let stream = Box::pin(tokio_stream::iter(events));
let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4);
let result = consume_model_cycle(stream, &event_tx, &CancellationToken::new())
.await
.unwrap();
assert_eq!(result.calls[0].name, "multi_tool_use_parallel");
assert!(matches!(
event_rx.recv().await,
Some(RunEvent::ToolCallStart { name, .. }) if name == "multi_tool_use_parallel"
));
}
#[tokio::test]
async fn usage_is_forwarded_before_the_provider_call_finishes() {
let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4);
+34 -17
View File
@@ -62,6 +62,40 @@ impl Store {
.await?)
}
pub async fn append_cursor_trace_request(
&self,
request_id: &str,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: &serde_json::Value,
) -> Result<()> {
let metadata_json = serde_json::to_string(metadata)?;
let blob_id = BlobId::digest(data);
let _write = self.writes.lock().await;
let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?;
Self::put_blob_tx(&mut tx, &blob_id, data, &[]).await?;
Self::link_cursor_trace_artifact_tx(
&mut tx,
request_id,
artifact_type,
source,
&blob_id,
&metadata_json,
)
.await?;
sqlx::query(
"UPDATE cursor_run_traces
SET request_bytes = request_bytes + ? WHERE request_id = ?",
)
.bind(as_i64(data.len()))
.bind(request_id)
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}
pub async fn append_cursor_trace_artifact(
&self,
request_id: &str,
@@ -144,23 +178,6 @@ impl Store {
Ok(())
}
pub async fn add_cursor_trace_request_bytes(
&self,
request_id: &str,
bytes: usize,
) -> Result<()> {
let _write = self.writes.lock().await;
sqlx::query(
"UPDATE cursor_run_traces
SET request_bytes = request_bytes + ? WHERE request_id = ?",
)
.bind(as_i64(bytes))
.bind(request_id)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> {
let now = now_ms();
let _write = self.writes.lock().await;
+99 -1
View File
@@ -8,6 +8,12 @@ use crate::{
use super::{now_ms, Store};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct ContextUsageAnchor {
pub(crate) context_input_tokens: u64,
pub(crate) message_count: usize,
}
#[derive(Clone, Debug)]
pub(crate) struct BufferedLlmChunk {
pub(crate) seq: i64,
@@ -266,6 +272,37 @@ impl Store {
Ok(())
}
pub(crate) async fn latest_context_usage(
&self,
conversation_id: &str,
) -> Result<Option<ContextUsageAnchor>> {
let row = sqlx::query(
"SELECT usage_json, message_count FROM llm_calls
WHERE conversation_id = ?
AND json_extract(usage_json, '$.context_input_tokens') IS NOT NULL
ORDER BY created_at_ms DESC, rowid DESC
LIMIT 1",
)
.bind(conversation_id)
.fetch_optional(&self.pool)
.await?;
let Some(row) = row else {
return Ok(None);
};
let usage: Usage = serde_json::from_str(row.try_get("usage_json")?)?;
let Some(context_input_tokens) = usage.context_input_tokens else {
return Ok(None);
};
let message_count = row.try_get::<i64, _>("message_count")?;
let Ok(message_count) = usize::try_from(message_count) else {
return Ok(None);
};
Ok(Some(ContextUsageAnchor {
context_input_tokens,
message_count,
}))
}
pub async fn llm_calls(&self, limit: i64) -> Result<Vec<LlmCallSummary>> {
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
.bind(limit.clamp(1, 500))
@@ -415,10 +452,71 @@ mod tests {
.await
.unwrap();
let overview = store
.overview(None, None, Some(&format!("[\"{plugin_model}\"]")))
.overview(None, None, Some(&format!("[\"{plugin_model}\"]")), None)
.await
.unwrap();
assert_eq!(overview.metrics.llm_calls, 1);
assert_eq!(overview.metrics.successful_calls, 1);
}
#[tokio::test]
async fn latest_context_usage_follows_conversation_chronology() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
for (call_id, model_id, context_input_tokens, message_count) in [
("call-a-1", "model-a", 100_u64, 3_usize),
("call-b", "model-b", 200_u64, 5_usize),
("call-a-2", "model-a", 300_u64, 7_usize),
] {
store
.start_llm_call(&NewLlmCall {
call_id: call_id.into(),
run_id: format!("run-{call_id}"),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: model_id.into(),
provider_type: ProviderType::Plugin,
provider_url: "plugin://test".into(),
request_type: ProviderType::Plugin,
request_url: "plugin://test".into(),
model_id: model_id.into(),
display_name: model_id.into(),
reasoning_effort: None,
fast: false,
message_count,
tool_count: 0,
detailed: false,
})
.await
.unwrap();
store
.record_llm_usage(
call_id,
Usage {
input_tokens: Some(context_input_tokens),
context_input_tokens: Some(context_input_tokens),
output_tokens: Some(10),
total_tokens: Some(context_input_tokens + 10),
..Default::default()
},
)
.await
.unwrap();
}
assert_eq!(
store.latest_context_usage("conversation").await.unwrap(),
Some(ContextUsageAnchor {
context_input_tokens: 300,
message_count: 7,
})
);
assert_eq!(store.latest_context_usage("other").await.unwrap(), None);
}
}
+1 -1
View File
@@ -19,7 +19,7 @@ mod writer;
pub use cas::*;
pub(crate) use cursor_traces::BufferedCursorTraceChunk;
pub(crate) use llm_calls::BufferedLlmChunk;
pub(crate) use llm_calls::{BufferedLlmChunk, ContextUsageAnchor};
pub use runs::*;
pub use settings::*;
pub(crate) use sqlite::now_ms;
+12 -12
View File
@@ -323,6 +323,18 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
})
}
fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result<Option<u64>> {
row.try_get::<Option<i64>, _>(column)?
.map(|value| {
u64::try_from(value).map_err(|_| Error::Config(format!("{column} cannot be negative")))
})
.transpose()
}
fn to_i64(value: u64) -> Result<i64> {
i64::try_from(value).map_err(|_| Error::Config("token value is too large".into()))
}
#[cfg(test)]
mod tests {
use super::*;
@@ -387,15 +399,3 @@ mod tests {
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| {
u64::try_from(value).map_err(|_| Error::Config(format!("{column} cannot be negative")))
})
.transpose()
}
fn to_i64(value: u64) -> Result<i64> {
i64::try_from(value).map_err(|_| Error::Config("token value is too large".into()))
}
+26 -8
View File
@@ -15,6 +15,7 @@ use super::Store;
const OVERVIEW_DAYS: u64 = 365;
const MAX_RANGE_BUCKETS: i64 = 60;
const MAX_EXPLICIT_BUCKETS: i64 = 1440;
const MINUTE_MS: i64 = 60_000;
const HOUR_MS: i64 = 60 * MINUTE_MS;
const DAY_MS: i64 = 24 * HOUR_MS;
@@ -25,6 +26,7 @@ impl Store {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<&str>,
bucket_ms: Option<i64>,
) -> Result<Overview> {
let call_row = sqlx::query(
"SELECT
@@ -84,7 +86,7 @@ impl Store {
};
let (token_usage_granularity, bucket_ms, series_start_ms, bucket_count) =
token_usage_buckets(start_ms, end_ms);
token_usage_buckets(start_ms, end_ms, bucket_ms);
let rows = sqlx::query(&format!(
"SELECT
(created_at_ms / {bucket_ms}) * {bucket_ms} AS bucket_start_ms,
@@ -146,20 +148,36 @@ impl Store {
fn token_usage_buckets(
start_ms: Option<i64>,
end_ms: Option<i64>,
requested_bucket_ms: Option<i64>,
) -> (TokenUsageGranularity, i64, i64, i64) {
if let (Some(start_ms), Some(end_ms)) = (start_ms, end_ms) {
let duration_ms = end_ms.saturating_sub(start_ms).max(1);
let (granularity, bucket_ms) = if duration_ms <= HOUR_MS {
(TokenUsageGranularity::Minute, MINUTE_MS)
} else if duration_ms <= MAX_RANGE_BUCKETS * HOUR_MS {
(TokenUsageGranularity::Hour, HOUR_MS)
let explicit_bucket_ms = requested_bucket_ms.filter(|bucket_ms| *bucket_ms >= MINUTE_MS);
let bucket_ms = explicit_bucket_ms.unwrap_or({
if duration_ms <= HOUR_MS {
MINUTE_MS
} else if duration_ms <= MAX_RANGE_BUCKETS * HOUR_MS {
HOUR_MS
} else {
DAY_MS
}
});
let granularity = if bucket_ms < HOUR_MS {
TokenUsageGranularity::Minute
} else if bucket_ms < DAY_MS {
TokenUsageGranularity::Hour
} else {
(TokenUsageGranularity::Day, DAY_MS)
TokenUsageGranularity::Day
};
let max_buckets = if explicit_bucket_ms.is_some() {
MAX_EXPLICIT_BUCKETS
} else {
MAX_RANGE_BUCKETS
};
let last_bucket_ms = end_ms.saturating_sub(1).div_euclid(bucket_ms) * bucket_ms;
let first_bucket_ms = start_ms.div_euclid(bucket_ms) * bucket_ms;
let bucket_count = ((last_bucket_ms - first_bucket_ms).div_euclid(bucket_ms) + 1)
.clamp(1, MAX_RANGE_BUCKETS);
let bucket_count =
((last_bucket_ms - first_bucket_ms).div_euclid(bucket_ms) + 1).clamp(1, max_buckets);
let series_start_ms =
last_bucket_ms.saturating_sub((bucket_count - 1).saturating_mul(bucket_ms));
return (granularity, bucket_ms, series_start_ms, bucket_count);
+137 -34
View File
@@ -1,6 +1,4 @@
//! Persists application settings.
use std::collections::HashSet;
use serde::{Deserialize, Serialize};
use crate::Result;
@@ -12,8 +10,12 @@ const PROXY_SETTINGS_KEY: &str = "outbound_proxy";
const TAB_SETTINGS_KEY: &str = "cursor_tab";
const INSTALLATION_ID_KEY: &str = "installation_id";
const DESKTOP_SETTINGS_KEY: &str = "desktop_lifecycle";
const DISABLED_PLUGIN_MODELS_KEY: &str = "disabled_plugin_models";
const DISABLED_PLUGIN_ACCOUNTS_KEY: &str = "disabled_plugin_accounts";
const COMMIT_SETTINGS_KEY: &str = "commit_settings";
const CURSOR_TAKEOVER_ENABLED_KEY: &str = "cursor_takeover_enabled";
/// Embedded default system prompts for commit message generation.
pub const DEFAULT_COMMIT_PROMPT_ZH_CN: &str = include_str!("../../prompt/cursor/commit/zh-CN.md");
pub const DEFAULT_COMMIT_PROMPT_EN_US: &str = include_str!("../../prompt/cursor/commit/en-US.md");
pub const PUBLIC_TAB_SERVICE_URL: &str = "https://tab.leokun.cn";
@@ -27,7 +29,7 @@ pub struct PortSettings {
#[serde(rename_all = "snake_case")]
pub enum ProxyMode {
#[default]
System,
Default,
Custom,
}
@@ -83,6 +85,54 @@ impl TabSettings {
}
}
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub enum CommitPromptLocale {
#[default]
#[serde(rename = "zh-CN")]
ZhCn,
#[serde(rename = "en-US")]
EnUs,
}
impl CommitPromptLocale {
pub fn default_prompt(self) -> &'static str {
match self {
Self::ZhCn => DEFAULT_COMMIT_PROMPT_ZH_CN.trim(),
Self::EnUs => DEFAULT_COMMIT_PROMPT_EN_US.trim(),
}
}
}
/// User preferences for Git commit message generation.
///
/// Empty `model_id` means 直连: forward the original Cursor RPC unchanged.
/// A non-empty value is the stable identifier of a configured built-in or
/// plugin model, and the request is generated locally through that model.
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub struct CommitSettings {
#[serde(default)]
pub model_id: String,
#[serde(default)]
pub prompt: String,
#[serde(default)]
pub prompt_locale: CommitPromptLocale,
}
impl CommitSettings {
pub fn is_direct(&self) -> bool {
self.model_id.trim().is_empty()
}
pub fn effective_prompt(&self) -> &str {
let trimmed = self.prompt.trim();
if trimmed.is_empty() {
self.prompt_locale.default_prompt()
} else {
trimmed
}
}
}
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub struct ProxySettingsInput {
pub mode: ProxyMode,
@@ -111,6 +161,32 @@ pub(crate) struct ProxySettingsSecret {
}
impl Store {
pub(crate) async fn cursor_takeover_enabled(&self) -> Result<bool> {
let value = sqlx::query_scalar::<_, String>(
"SELECT value_json FROM service_settings WHERE setting_key = ?",
)
.bind(CURSOR_TAKEOVER_ENABLED_KEY)
.fetch_optional(&self.pool)
.await?;
value
.map(|value| serde_json::from_str(&value).map_err(Into::into))
.unwrap_or(Ok(true))
}
pub(crate) async fn set_cursor_takeover_enabled(&self, enabled: bool) -> Result<()> {
let value_json = serde_json::to_string(&enabled)?;
let _write = self.writes.lock().await;
sqlx::query(
"INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms",
)
.bind(CURSOR_TAKEOVER_ENABLED_KEY)
.bind(value_json)
.bind(now_ms())
.execute(&self.pool)
.await?;
Ok(())
}
pub(crate) async fn installation_id(&self) -> Result<String> {
let generated = uuid::Uuid::new_v4().to_string();
let _write = self.writes.lock().await;
@@ -304,55 +380,82 @@ impl Store {
Ok(())
}
pub async fn disabled_plugin_models(&self) -> Result<HashSet<String>> {
pub async fn commit_settings(&self) -> Result<CommitSettings> {
let value = sqlx::query_scalar::<_, String>(
"SELECT value_json FROM service_settings WHERE setting_key = ?",
)
.bind(DISABLED_PLUGIN_MODELS_KEY)
.bind(COMMIT_SETTINGS_KEY)
.fetch_optional(&self.pool)
.await?;
value
.map(|value| serde_json::from_str(&value).map_err(Into::into))
.unwrap_or_else(|| Ok(HashSet::new()))
.unwrap_or_else(|| Ok(CommitSettings::default()))
}
pub async fn set_disabled_plugin_models(&self, model_ids: &HashSet<String>) -> Result<()> {
let value_json = serde_json::to_string(model_ids)?;
pub async fn set_commit_settings(&self, settings: CommitSettings) -> Result<CommitSettings> {
let settings = CommitSettings {
model_id: settings.model_id.trim().to_owned(),
prompt: settings.prompt.trim().to_owned(),
prompt_locale: settings.prompt_locale,
};
let value_json = serde_json::to_string(&settings)?;
let _write = self.writes.lock().await;
sqlx::query(
"INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms",
)
.bind(DISABLED_PLUGIN_MODELS_KEY)
.bind(COMMIT_SETTINGS_KEY)
.bind(value_json)
.bind(now_ms())
.execute(&self.pool)
.await?;
Ok(())
Ok(settings)
}
}
pub async fn disabled_plugin_accounts(&self) -> Result<HashSet<String>> {
let value = sqlx::query_scalar::<_, String>(
"SELECT value_json FROM service_settings WHERE setting_key = ?",
)
.bind(DISABLED_PLUGIN_ACCOUNTS_KEY)
.fetch_optional(&self.pool)
.await?;
value
.map(|value| serde_json::from_str(&value).map_err(Into::into))
.unwrap_or_else(|| Ok(HashSet::new()))
#[cfg(test)]
mod tests {
use super::{
CommitPromptLocale, CommitSettings, ProxyMode, DEFAULT_COMMIT_PROMPT_EN_US,
DEFAULT_COMMIT_PROMPT_ZH_CN,
};
#[test]
fn default_commit_prompt_follows_its_saved_locale() {
for (prompt_locale, expected) in [
(CommitPromptLocale::ZhCn, DEFAULT_COMMIT_PROMPT_ZH_CN),
(CommitPromptLocale::EnUs, DEFAULT_COMMIT_PROMPT_EN_US),
] {
let settings = CommitSettings {
prompt_locale,
..CommitSettings::default()
};
assert_eq!(settings.effective_prompt(), expected.trim());
}
}
#[test]
fn custom_commit_prompt_does_not_change_with_locale() {
for prompt_locale in [CommitPromptLocale::ZhCn, CommitPromptLocale::EnUs] {
let settings = CommitSettings {
prompt: "custom prompt".into(),
prompt_locale,
..CommitSettings::default()
};
assert_eq!(settings.effective_prompt(), "custom prompt");
}
}
pub async fn set_disabled_plugin_accounts(&self, account_ids: &HashSet<String>) -> Result<()> {
let value_json = serde_json::to_string(account_ids)?;
let _write = self.writes.lock().await;
sqlx::query(
"INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms",
)
.bind(DISABLED_PLUGIN_ACCOUNTS_KEY)
.bind(value_json)
.bind(now_ms())
.execute(&self.pool)
.await?;
Ok(())
#[test]
fn default_proxy_mode_uses_the_default_wire_value() {
assert_eq!(ProxyMode::default(), ProxyMode::Default);
assert_eq!(
serde_json::to_string(&ProxyMode::default()).unwrap(),
"\"default\""
);
assert_eq!(
serde_json::from_str::<ProxyMode>("\"default\"").unwrap(),
ProxyMode::Default
);
assert!(serde_json::from_str::<ProxyMode>("\"system\"").is_err());
}
}
+251
View File
@@ -0,0 +1,251 @@
//! Verifies the local WriteGitCommitMessage engine end to end on the wire.
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use axum::{
body::{to_bytes, Body},
http::{header, Request, StatusCode},
};
use cursor_server::{
api::cursor,
cursor::{
prompting::{PromptAssets, PromptCompiler},
protocol::{connect, proto::aiserver::v1 as ai},
transport::TransportRegistry,
},
model::{ContentPart, ModelConfigInput, ModelType, ProjectedContent, OPENAI_CHAT_ENDPOINT},
network::NetworkClients,
provider::{FinishReason, ModelEvent},
store::{CommitPromptLocale, CommitSettings, DEFAULT_COMMIT_PROMPT_ZH_CN},
};
use tower::ServiceExt;
async fn commit_router(
store: cursor_server::store::Store,
provider: fake_provider::FakeProvider,
) -> axum::Router {
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let clients = NetworkClients::new(store.clone());
let registry = TransportRegistry::new(store, Arc::new(provider), PromptCompiler::new(assets));
cursor::router(registry, clients).unwrap()
}
fn model_input(model_id: &str) -> ModelConfigInput {
ModelConfigInput {
sort_order: 1,
display_name: "Qwen Flash".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1".into(),
use_full_url: false,
api_key: "test-key".into(),
tooltip_data: "模型介绍".into(),
model_id: model_id.into(),
reasoning_effort: None,
openai_endpoint: 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,
}
}
async fn post_commit_message(
router: axum::Router,
request: ai::WriteGitCommitMessageRequest,
) -> axum::response::Response {
let body = connect::encode_message(&request).unwrap();
router
.oneshot(
Request::post("/aiserver.v1.AiService/WriteGitCommitMessage")
.header(header::CONTENT_TYPE, "application/proto")
.body(Body::from(body))
.unwrap(),
)
.await
.unwrap()
}
fn diff_request(diff: &str) -> ai::WriteGitCommitMessageRequest {
ai::WriteGitCommitMessageRequest {
diffs: vec![diff.into()],
previous_commit_messages: vec!["feat: 上一次提交".into()],
explicit_context: None,
}
}
#[tokio::test]
async fn commit_message_is_generated_through_configured_model() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash.clone(),
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::TextStart,
ModelEvent::TextDelta("```\nCommit message: feat: 新增提交引擎\n```".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let router = commit_router(store, provider.clone()).await;
let response = post_commit_message(router, diff_request("diff --git a/engine.rs")).await;
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 64 * 1024).await.unwrap();
let decoded: ai::WriteGitCommitMessageResponse = prost::Message::decode(&body[..]).unwrap();
assert_eq!(decoded.commit_message, "feat: 新增提交引擎");
let requests = provider.requests();
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].prompt.instructions,
DEFAULT_COMMIT_PROMPT_ZH_CN.trim()
);
let ProjectedContent::Parts(parts) = &requests[0].history[0].content else {
panic!("expected user text parts");
};
let ContentPart::Text { text } = &parts[0] else {
panic!("expected text part");
};
assert!(text.contains("diff --git a/engine.rs"));
assert!(text.contains("- feat: 上一次提交"));
}
#[tokio::test]
async fn custom_prompt_and_model_from_commit_settings_are_used() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-coder"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: "自定义提交提示词".into(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::TextDelta("chore: 清理旧代码".into()),
ModelEvent::Done(FinishReason::Stop),
]);
let router = commit_router(store, provider.clone()).await;
let response = post_commit_message(router, diff_request("diff --git a/old.rs")).await;
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 64 * 1024).await.unwrap();
let decoded: ai::WriteGitCommitMessageResponse = prost::Message::decode(&body[..]).unwrap();
assert_eq!(decoded.commit_message, "chore: 清理旧代码");
assert_eq!(
provider.requests()[0].prompt.instructions,
"自定义提交提示词"
);
}
#[tokio::test]
async fn empty_diffs_are_rejected_when_generating() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
let router = commit_router(store, provider).await;
let response = post_commit_message(router, ai::WriteGitCommitMessageRequest::default()).await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("diffs are required"));
}
#[tokio::test]
async fn tool_call_events_are_rejected() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![ModelEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "shell".into(),
}]);
let router = commit_router(store, provider).await;
let response = post_commit_message(router, diff_request("diff --git a/x.rs")).await;
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("must not invoke tools"));
}
#[tokio::test]
async fn unconfigured_model_is_rejected() {
let (_directory, store) = fixtures::temp_store().await;
store
.set_commit_settings(CommitSettings {
model_id: "missing-hash".into(),
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
let router = commit_router(store, provider).await;
let response = post_commit_message(router, diff_request("diff --git a/x.rs")).await;
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("missing-hash"));
}
+123 -3
View File
@@ -267,6 +267,16 @@ async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_toke
.await;
assert_eq!(second.summary_started, 1);
assert_eq!(second.summary_completed, 1);
assert_eq!(
&second.interaction_events[..4],
&[
"token_delta:0",
"summary_started",
"summary_completed",
"token_delta:0",
],
"automatic compaction must publish estimated usage before summarizing and zero usage after"
);
let compacted_tokens = second
.checkpoints
.iter()
@@ -297,6 +307,108 @@ async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_toke
}));
}
#[tokio::test]
async fn incremental_preflight_uses_conversation_anchor_across_model_switch() {
let (_directory, store) = fixtures::temp_store().await;
let model_a = store
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Anchor Model A".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Anchor Model A".into(),
model_id: "anchor-model-a".into(),
reasoning_effort: None,
openai_endpoint: 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,
})
.await
.unwrap();
let model_b = store
.create_model(&ModelConfigInput {
sort_order: 1,
display_name: "Anchor Model B".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Anchor Model B".into(),
model_id: "anchor-model-b".into(),
reasoning_effort: None,
openai_endpoint: 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: Some(200_000),
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(text_response("old answer", 103_904, 12));
provider.push(text_response("new answer", 104_000, 12));
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 first = run(
&registry,
"anchor-first",
user_request(
"anchor-conversation",
"anchor-user-1",
&"x".repeat(400_000),
&model_a.model_hash,
None,
),
)
.await;
let second = run(
&registry,
"anchor-second",
user_request(
"anchor-conversation",
"anchor-user-2",
"short follow-up",
&model_b.model_hash,
first.checkpoints.last().cloned(),
),
)
.await;
assert_eq!(second.summary_started, 0);
assert_eq!(second.summary_completed, 0);
assert_eq!(provider.requests().len(), 2);
}
#[tokio::test]
async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() {
let (_directory, store) = fixtures::temp_store().await;
@@ -367,6 +479,7 @@ struct Output {
summary_completed: usize,
turn_ended: usize,
token_delta: usize,
interaction_events: Vec<String>,
}
async fn run(
@@ -415,16 +528,23 @@ async fn run(
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
match update.message {
Some(pb::interaction_update::Message::SummaryStarted(_)) => {
output.summary_started += 1
output.summary_started += 1;
output.interaction_events.push("summary_started".into());
}
Some(pb::interaction_update::Message::Summary(delta)) => {
output.summary.push_str(&delta.summary)
}
Some(pb::interaction_update::Message::SummaryCompleted(_)) => {
output.summary_completed += 1
output.summary_completed += 1;
output.interaction_events.push("summary_completed".into());
}
Some(pb::interaction_update::Message::TurnEnded(_)) => output.turn_ended += 1,
Some(pb::interaction_update::Message::TokenDelta(_)) => output.token_delta += 1,
Some(pb::interaction_update::Message::TokenDelta(delta)) => {
output.token_delta += 1;
output
.interaction_events
.push(format!("token_delta:{}", delta.tokens));
}
_ => {}
}
}
+3 -1
View File
@@ -21,6 +21,7 @@ use cursor_server::{
proto::{agent::v1 as pb, aiserver::v1 as ai},
},
cursor::transport::TransportRegistry,
network::NetworkClients,
};
use flate2::{write::GzEncoder, Compression};
use prost::Message;
@@ -102,6 +103,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() {
.as_path(),
)
.unwrap();
let clients = NetworkClients::new(store.clone());
let registry = TransportRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
@@ -118,7 +120,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() {
encoder.write_all(&wire).unwrap();
let compressed = encoder.finish().unwrap();
let response = cursor::router(registry)
let response = cursor::router(registry, clients)
.unwrap()
.oneshot(
Request::post("/aiserver.v1.BidiService/BidiAppend")
+129
View File
@@ -0,0 +1,129 @@
//! Verifies that Cursor trace persistence is ordered and detached from producers.
#[path = "support/fixtures.rs"]
mod fixtures;
use std::time::{Duration, Instant};
use bytes::Bytes;
use cursor_server::{cursor::services::observability::CursorTraceService, store::Store};
use sqlx::{Connection, SqliteConnection};
#[tokio::test]
async fn trace_producers_do_not_wait_for_sqlite_and_artifacts_stay_ordered() {
let directory = tempfile::tempdir().unwrap();
let url = format!("sqlite://{}", directory.path().join("test.db").display());
let store = Store::connect(&url).await.unwrap();
store.set_detailed_logging(true).await.unwrap();
let traces = CursorTraceService::new(store.clone());
let recorder = traces.recorder("trace-queue-order");
recorder.begin(Some("conversation-1"), "local_byok", Some("model-1"));
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if store
.cursor_trace("trace-queue-order")
.await
.unwrap()
.is_some()
{
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
let mut write_lock = SqliteConnection::connect(&url).await.unwrap();
sqlx::query("BEGIN IMMEDIATE")
.execute(&mut write_lock)
.await
.unwrap();
recorder.request(
"bidi_request",
Bytes::from_static(b"request-0"),
serde_json::json!({
"append_seqno": 0,
"accepted": true,
"route_outcome": "local"
}),
);
tokio::time::sleep(Duration::from_millis(25)).await;
let started = Instant::now();
let mut seqnos = (1..64).collect::<Vec<_>>();
for pair in seqnos.chunks_mut(2) {
pair.reverse();
}
for seqno in seqnos {
recorder.request(
"bidi_request",
Bytes::from(format!("request-{seqno}")),
serde_json::json!({
"append_seqno": seqno,
"accepted": true,
"route_outcome": "local"
}),
);
}
recorder.finish(None);
assert!(started.elapsed() < Duration::from_millis(100));
sqlx::query("ROLLBACK")
.execute(&mut write_lock)
.await
.unwrap();
let artifacts = tokio::time::timeout(Duration::from_secs(5), async {
loop {
let artifacts = store
.cursor_trace_artifacts("trace-queue-order")
.await
.unwrap();
if artifacts.len() == 64 {
break artifacts;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
for (expected, artifact) in artifacts.iter().enumerate() {
assert_eq!(artifact.seq, expected as i64);
assert_eq!(artifact.metadata["append_seqno"], expected as i64);
}
let trace = store
.cursor_trace("trace-queue-order")
.await
.unwrap()
.unwrap();
assert_eq!(trace.status, "completed");
assert_eq!(
trace.request_bytes,
(0..64)
.map(|seqno| format!("request-{seqno}").len() as i64)
.sum::<i64>()
);
}
#[tokio::test]
async fn events_for_disabled_detailed_logging_are_discarded_off_path() {
let (_directory, store) = fixtures::temp_store().await;
let traces = CursorTraceService::new(store.clone());
let recorder = traces.recorder("trace-disabled");
recorder.begin(None, "local_byok", Some("model-1"));
recorder.request(
"bidi_request",
Bytes::from_static(b"body"),
serde_json::json!({"append_seqno": 0}),
);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(store
.cursor_trace("trace-disabled")
.await
.unwrap()
.is_none());
}
@@ -0,0 +1,72 @@
//! Verifies registry ownership follows the transport actor rather than output subscriptions.
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{sync::Arc, time::Duration};
use cursor_server::cursor::{
conversation::TransportCommand,
prompting::{PromptAssets, PromptCompiler},
transport::TransportRegistry,
};
async fn registry() -> (tempfile::TempDir, TransportRegistry) {
let (directory, store) = fixtures::temp_store().await;
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
(
directory,
TransportRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
),
)
}
#[tokio::test]
async fn actor_exit_removes_the_matching_transport_and_allows_a_new_generation() {
let (_directory, registry) = registry().await;
let first = registry.get_or_create("lifecycle-request").await.unwrap();
assert!(registry.local("lifecycle-request").await.is_some());
first.command(TransportCommand::Disconnect).await.unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if registry.local("lifecycle-request").await.is_none() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
let second = registry.get_or_create("lifecycle-request").await.unwrap();
assert_eq!(second.request_id(), "lifecycle-request");
assert!(registry.local("lifecycle-request").await.is_some());
second.command(TransportCommand::Disconnect).await.unwrap();
}
#[tokio::test]
async fn dropping_an_output_subscription_does_not_remove_the_transport() {
let (_directory, registry) = registry().await;
let handle = registry
.get_or_create("subscription-request")
.await
.unwrap();
let subscription = handle.subscribe();
drop(subscription);
tokio::time::sleep(Duration::from_millis(25)).await;
assert!(registry.local("subscription-request").await.is_some());
handle.command(TransportCommand::Disconnect).await.unwrap();
}
+128
View File
@@ -422,6 +422,85 @@ async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() {
assert_eq!(output.recv().await, None);
}
#[tokio::test]
async fn queued_user_message_after_turn_ended_starts_the_next_turn() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(text_response("first turn"));
provider.push(text_response("queued turn"));
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("queued-after-turn").await.unwrap();
let mut output = handle.subscribe();
handle
.command(TransportCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"queued-after-turn",
"queued-after-turn-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
wait_for_turn_ended(&handle, &mut output, &mut append_seqno).await;
assert_transport_remains_open(&handle, &mut output, &mut append_seqno).await;
cursor_server::api::cursor::bidi::append(
&registry,
cursor_server::api::cursor::bidi::DecodedAppend {
request_id: "queued-after-turn".into(),
seqno: append_seqno,
message: runtime_user_message(),
},
None,
)
.await
.unwrap();
append_seqno += 1;
let mut text = String::new();
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("queued turn closed without EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
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 {
text.push_str(&delta.text);
}
}
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
assert!(text.contains("queued turn"));
let requests = provider.requests();
assert_eq!(requests.len(), 2);
assert_eq!(
&requests[1].history[..requests[0].history.len()],
requests[0].history.as_slice(),
"queued continuation must preserve the first provider request as a prefix"
);
let history = serde_json::to_string(&requests[1].history).unwrap();
assert!(history.contains("queued follow-up"));
}
#[tokio::test]
async fn runtime_user_message_action_interrupts_and_continues_with_new_message() {
let (_directory, store) = fixtures::temp_store().await;
@@ -1732,6 +1811,55 @@ async fn run_to_end(
}
}
async fn wait_for_turn_ended(
handle: &cursor_server::cursor::TransportHandle,
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
append_seqno: &mut i64,
) {
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before turnEnded");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
let server = pb::AgentServerMessage::decode(payload).unwrap();
let turn_ended = matches!(
server.message,
Some(pb::agent_server_message::Message::InteractionUpdate(
pb::InteractionUpdate {
message: Some(pb::interaction_update::Message::TurnEnded(_)),
}
))
);
acknowledge_kv(handle, append_seqno, &frame).await;
if turn_ended {
return;
}
}
}
async fn assert_transport_remains_open(
handle: &cursor_server::cursor::TransportHandle,
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
append_seqno: &mut i64,
) {
let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(100);
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
let Ok(Some(frame)) = tokio::time::timeout(remaining, output.recv()).await else {
return;
};
let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(
flags & connect::END_STREAM_FLAG,
0,
"turnEnded closed the transport before the queued action arrived"
);
acknowledge_kv(handle, append_seqno, &frame).await;
}
}
async fn acknowledge_kv(
handle: &cursor_server::cursor::TransportHandle,
append_seqno: &mut i64,
+6 -3
View File
@@ -106,7 +106,7 @@ async fn decode<M: Message + Default>(response: Response<Body>) -> M {
#[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 upstream = CursorProxy::cursor(cursor_server::network::NetworkClients::new(store));
let rules_dir = tempfile::tempdir().unwrap();
let rules_root = rules_dir.path().join("rules");
let service = KnowledgeService::with_root(rules_root.clone()).unwrap();
@@ -125,7 +125,10 @@ async fn offline_crud_round_trip_persists_markdown() {
.unwrap();
let added: AddResponse = decode(response).await;
assert!(added.success);
assert!(added.id.starts_with("local-"), "offline add uses a local id");
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(),
@@ -193,7 +196,7 @@ async fn offline_crud_round_trip_persists_markdown() {
#[tokio::test]
async fn updating_missing_rule_reports_failure() {
let (_store_dir, store) = fixtures::temp_store().await;
let upstream = CursorProxy::cursor(store).unwrap();
let upstream = CursorProxy::cursor(cursor_server::network::NetworkClients::new(store));
let rules_dir = tempfile::tempdir().unwrap();
let service = KnowledgeService::with_root(rules_dir.path().join("rules")).unwrap();
+31 -5
View File
@@ -16,6 +16,7 @@ use cursor_server::{
model::{ContentPart, ProjectedContent},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
#[tokio::test]
async fn local_markdown_rules_land_in_the_request_context_message() {
@@ -58,18 +59,30 @@ async fn local_markdown_rules_land_in_the_request_context_message() {
.await
.unwrap();
let mut append_seqno = 1;
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 {
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break;
}
// The Run waits for the client to confirm every conversation Blob write,
// so the stream only advances once each KvServerMessage is acknowledged.
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) =
pb::AgentServerMessage::decode(payload).unwrap().message
{
handle
.command(TransportCommand::Append {
seqno: append_seqno,
message: Box::new(set_blob_result(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
}
let requests = provider.requests();
@@ -102,6 +115,19 @@ async fn local_markdown_rules_land_in_the_request_context_message() {
registry.shutdown().await;
}
fn set_blob_result(id: u32) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::KvClientMessage(
pb::KvClientMessage {
id,
message: Some(pb::kv_client_message::Message::SetBlobResult(
pb::SetBlobResult { error: None },
)),
},
)),
}
}
fn user_run() -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
+145
View File
@@ -1253,6 +1253,151 @@ fn client_run_for_model(
}
}
/// A provider that reuses a tool call id across two rounds of the same run
/// must not wedge the run. `ToolDispatcher::start_batch` skips any call whose
/// id is already in `ToolBatchState::completed`, and that set is never cleared
/// for the lifetime of the run, so the skipped call produced no completion and
/// `tool_round::execute` waited forever for a result that could never arrive.
#[tokio::test]
async fn duplicate_tool_call_id_across_rounds_does_not_wedge_the_run() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
for _ in 0..2 {
provider.push(vec![
ModelEvent::Start {
model_call_id: "ignored".into(),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":\"/tmp/a\"}".into(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]);
}
provider.push(vec![
ModelEvent::Start {
model_call_id: "ignored".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("done".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 registry = TransportRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
);
let handle = registry.get_or_create("tool-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(TransportCommand::Append {
seqno: 0,
message: Box::new(client_run()),
})
.await
.unwrap();
let mut seqno = 1;
let mut execs = 0;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(10), output.recv())
.await
.expect("run must not hang on a reused tool call id")
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(TransportCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
execs += 1;
let exec_id = exec.id;
handle
.command(TransportCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id: exec_id,
exec_id: String::new(),
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(
pb::ReadSuccess {
path: "/tmp/a".into(),
total_lines: 1,
file_size: 1,
output: Some(
pb::read_success::Output::Content(
"x".into(),
),
),
..Default::default()
},
)),
},
)),
..Default::default()
},
)),
}),
})
.await
.unwrap();
seqno += 1;
handle
.command(TransportCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(
pb::agent_client_message::Message::ExecClientControlMessage(
pb::ExecClientControlMessage {
message: Some(
pb::exec_client_control_message::Message::StreamClose(
pb::ExecClientStreamClose { id: exec_id },
),
),
},
),
),
}),
})
.await
.unwrap();
seqno += 1;
}
_ => {}
}
}
// Both rounds have to reach the client, and the run has to get far enough
// to ask the provider a third time and finish.
assert_eq!(execs, 2);
assert_eq!(provider.requests().len(), 3);
}
fn client_run() -> pb::AgentClientMessage {
let user = pb::UserMessage {
text: "read it".into(),