mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 19:31:28 +08:00
Merge main and host authorization-code OAuth in core
This commit is contained in:
@@ -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);
|
||||
|
||||
@@ -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}`);
|
||||
|
||||
@@ -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`
|
||||
@@ -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(<具体修改项>): 修改美术资源`
|
||||
@@ -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)?;
|
||||
}
|
||||
|
||||
@@ -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(®istry, &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(®istry, 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(®istry, 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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
@@ -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
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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?,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -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?))
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
))))
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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, ¤t, 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, ¤t, 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(
|
||||
®istry,
|
||||
&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(
|
||||
®istry,
|
||||
&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;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")]
|
||||
|
||||
@@ -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"))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(®istry, upstream, request).await;
|
||||
}
|
||||
generate_local(®istry, 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\"}");
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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(®istry).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(®istry).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(®istry).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"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
@@ -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}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,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,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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(""));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)))
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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(())
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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('&', "&")
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
.replace('"', """)
|
||||
.replace('\'', "'")
|
||||
}
|
||||
|
||||
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("<plugin>"));
|
||||
assert!(page.contains("Accounts & 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
@@ -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 {
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let request = request.timeout(Duration::from_secs(60));
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if 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, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
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", ¶ms).await.unwrap();
|
||||
let (_, _, second_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, body) = recorded_network_request(¶ms).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(¶ms).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", ¶ms).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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ mod normalize;
|
||||
mod openai_chat;
|
||||
mod openai_responses;
|
||||
mod recorder;
|
||||
mod request_template;
|
||||
mod router;
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
})]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -129,6 +129,7 @@ pub enum RunEvent {
|
||||
ToolCallEnd {
|
||||
index: usize,
|
||||
},
|
||||
UsageSnapshot(Usage),
|
||||
Usage(Usage),
|
||||
ExecuteToolRound {
|
||||
round_id: ToolRoundId,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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()))
|
||||
}
|
||||
|
||||
@@ -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
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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(
|
||||
®istry,
|
||||
"anchor-first",
|
||||
user_request(
|
||||
"anchor-conversation",
|
||||
"anchor-user-1",
|
||||
&"x".repeat(400_000),
|
||||
&model_a.model_hash,
|
||||
None,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let second = run(
|
||||
®istry,
|
||||
"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));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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(
|
||||
®istry,
|
||||
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,
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user