diff --git a/apps/desktop/src/features/plugins/PluginResourcePanels.tsx b/apps/desktop/src/features/plugins/PluginResourcePanels.tsx index 2c4f902..c94e69b 100644 --- a/apps/desktop/src/features/plugins/PluginResourcePanels.tsx +++ b/apps/desktop/src/features/plugins/PluginResourcePanels.tsx @@ -109,7 +109,7 @@ function OAuthMethodCard({ pluginId, resourceType, method, onConfigured }: { const next = await api.pluginOAuthBegin(pluginId, resourceType, method.id); setBegun(next); setStatus("polling"); - await api.copyCursorText(next.userCode).catch(() => undefined); + if (next.userCode) await api.copyCursorText(next.userCode).catch(() => undefined); await api.openExternalUrl(next.verificationUrlComplete || next.verificationUrl); } catch (cause) { setStatus("error"); @@ -117,13 +117,14 @@ function OAuthMethodCard({ pluginId, resourceType, method, onConfigured }: { } }; + const userCode = begun?.userCode; return {pluginText(method.displayName, locale)} {method.description && {pluginText(method.description, locale)}} - {begun && status === "polling" &&
+ {userCode && status === "polling" &&
{t("设备验证码")} - - +
} diff --git a/apps/desktop/src/shared/api.ts b/apps/desktop/src/shared/api.ts index c11c11c..2c2d5cf 100644 --- a/apps/desktop/src/shared/api.ts +++ b/apps/desktop/src/shared/api.ts @@ -213,10 +213,11 @@ export interface PluginResourceView { } export interface PluginAddMethod { - type: "oauth2.0"; + type: "oauth2.0" | "oauth2.authorization-code"; id: string; displayName: PluginLocalizedText; description: PluginLocalizedText | null; + callback?: { port: number | null; path: string | null }; } export interface PluginImportDescriptor { @@ -274,7 +275,7 @@ export interface PluginDescriptor { export interface PluginOAuthBegin { sessionId: string; - userCode: string; + userCode: string | null; verificationUrl: string; verificationUrlComplete: string | null; expiresAtMs: number; diff --git a/server/plugins/build-in/antigravity-auth/assets/antigravity.svg b/server/plugins/build-in/antigravity-auth/assets/antigravity.svg new file mode 100644 index 0000000..3ed10ab --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/assets/antigravity.svg @@ -0,0 +1 @@ +Antigravity \ No newline at end of file diff --git a/server/plugins/build-in/antigravity-auth/deno.json b/server/plugins/build-in/antigravity-auth/deno.json new file mode 100644 index 0000000..625b8d8 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/deno.json @@ -0,0 +1,13 @@ +{ + "imports": { + "cursor-byok:plugin": "../../../src/plugin/sdk/plugin.ts", + "cursor-byok:provider": "../../../src/plugin/sdk/provider.ts", + "cursor-byok:model": "../../../src/plugin/sdk/model.ts", + "cursor-byok:resource": "../../../src/plugin/sdk/resource.ts", + "cursor-byok:protocol/openai-chat": "../../../src/plugin/sdk/protocol/openai_chat.ts" + }, + "fmt": { + "lineWidth": 100, + "exclude": ["assets"] + } +} diff --git a/server/plugins/build-in/antigravity-auth/google_oauth.ts b/server/plugins/build-in/antigravity-auth/google_oauth.ts new file mode 100644 index 0000000..a23dfc3 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/google_oauth.ts @@ -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"; diff --git a/server/plugins/build-in/antigravity-auth/main.ts b/server/plugins/build-in/antigravity-auth/main.ts new file mode 100644 index 0000000..d7aa4be --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/main.ts @@ -0,0 +1,19 @@ +import { defineProviderPlugin } from "cursor-byok:plugin"; +import { antigravityAuthorizationCodeOAuth } from "./oauth.ts"; +import { antigravityProvider } from "./provider.ts"; +import { credentialImport, presentAccount, refreshAccount, RESOURCE_TYPE } from "./resources.ts"; + +export default defineProviderPlugin({ + providers: [antigravityProvider], + resources: [{ + type: RESOURCE_TYPE, + displayName: { + "en-US": "Google accounts & API keys", + "zh-CN": "Google 账号与 API 密钥", + }, + add: [antigravityAuthorizationCodeOAuth], + import: credentialImport, + present: presentAccount, + refresh: refreshAccount, + }], +}); diff --git a/server/plugins/build-in/antigravity-auth/models.ts b/server/plugins/build-in/antigravity-auth/models.ts new file mode 100644 index 0000000..9a148f6 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/models.ts @@ -0,0 +1,349 @@ +import type { JsonValue } from "cursor-byok:plugin"; +import type { ModelDefinition, ModelSnapshot, ModelSupport } from "cursor-byok:model"; +import { accountData } from "./resources.ts"; + +export const ANTIGRAVITY_SANDBOX_ENDPOINT = "https://daily-cloudcode-pa.sandbox.googleapis.com"; +export const ANTIGRAVITY_DAILY_ENDPOINT = "https://daily-cloudcode-pa.googleapis.com"; +export const ANTIGRAVITY_PROD_ENDPOINT = "https://cloudcode-pa.googleapis.com"; + +export const ANTIGRAVITY_ENDPOINTS = [ + ANTIGRAVITY_SANDBOX_ENDPOINT, + ANTIGRAVITY_DAILY_ENDPOINT, + ANTIGRAVITY_PROD_ENDPOINT, +]; + +const FETCH_AVAILABLE_MODELS_PATH = "/v1internal:fetchAvailableModels"; +export const ANTIGRAVITY_USER_AGENT = + "Antigravity/4.3.0 (Macintosh; Intel Mac OS X 10_15_7) Chrome/132.0.6834.160 Electron/39.2.3"; + +export const ANTIGRAVITY_CLIENT_HEADERS: Record = { + "x-client-name": "antigravity", + "x-client-version": "4.3.0", +}; + +const ANTIGRAVITY_DENYLIST = new Set(["chat_20706", "chat_23310"]); + +export const STATIC_ANTIGRAVITY_MODELS: ModelDefinition[] = [ + // Gemini 3.7 Series + { + id: "gemini-3.7-flash", + displayName: "Gemini 3.7 Flash", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "gemini-3.7-flash-high", + displayName: "Gemini 3.7 Flash (High)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.7-flash-medium", + displayName: "Gemini 3.7 Flash (Medium)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.7-flash-low", + displayName: "Gemini 3.7 Flash (Low)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.7-flash-tiered", + displayName: "Gemini 3.7 Flash (Tiered)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.7-flash-thinking", + displayName: "Gemini 3.7 Flash (Thinking)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + + // Gemini 3.6 Series + { + id: "gemini-3.6-flash-high", + displayName: "Gemini 3.6 Flash (High)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.6-flash-medium", + displayName: "Gemini 3.6 Flash (Medium)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.6-flash-low", + displayName: "Gemini 3.6 Flash (Low)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + + // Gemini 3.1 Pro Series + { + id: "gemini-3.1-pro-preview", + displayName: "Gemini 3.1 Pro Preview", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "gemini-3.1-pro-high", + displayName: "Gemini 3.1 Pro (High)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.1-pro-medium", + displayName: "Gemini 3.1 Pro (Medium)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.1-pro-low", + displayName: "Gemini 3.1 Pro (Low)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + + // Gemini 2.5 / 2.0 Series + { + id: "gemini-2.5-pro", + displayName: "Gemini 2.5 Pro", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "gemini-2.5-flash", + displayName: "Gemini 2.5 Flash", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "gemini-2.5-flash-thinking", + displayName: "Gemini 2.5 Flash Thinking", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-2.5-flash-lite", + displayName: "Gemini 2.5 Flash Lite", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-2.0-flash", + displayName: "Gemini 2.0 Flash", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-2.0-flash-lite", + displayName: "Gemini 2.0 Flash Lite", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + + // Claude Series (via Antigravity) + { + id: "claude-sonnet-4-6", + displayName: "Claude Sonnet 4.6 (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "claude-sonnet-4-6-thinking", + displayName: "Claude Sonnet 4.6 Thinking (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: [] }, + }, + { + id: "claude-opus-4-6-thinking", + displayName: "Claude 3.7 Opus Thinking (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "claude-3-7-sonnet", + displayName: "Claude 3.7 Sonnet (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "claude-3-5-sonnet", + displayName: "Claude 3.5 Sonnet (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "claude-3-5-haiku", + displayName: "Claude 3.5 Haiku (Antigravity)", + capabilities: { images: true }, + maxOutputTokens: 64000, + privateData: { reasoningEfforts: [] }, + }, + + // Other Models + { + id: "gpt-4o", + displayName: "GPT-4o (Antigravity / Gemini)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: ["low", "medium", "high"] }, + }, + { + id: "gpt-4o-mini", + displayName: "GPT-4o Mini (Antigravity / Gemini)", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gpt-oss-120b-medium", + displayName: "GPT OSS 120B Medium", + capabilities: { images: false }, + maxOutputTokens: 32768, + privateData: { reasoningEfforts: [] }, + }, + { + id: "gemini-3.1-flash-image", + displayName: "Gemini 3.1 Flash Image", + capabilities: { images: true }, + maxOutputTokens: 65536, + privateData: { reasoningEfforts: [] }, + }, +]; + +function object(value: unknown): Record | null { + return value !== null && typeof value === "object" && !Array.isArray(value) + ? value as Record + : null; +} + +function text(value: unknown): string | null { + return typeof value === "string" && value.trim() ? value.trim() : null; +} + +export function parseAntigravityModels(payload: unknown): ModelDefinition[] { + const root = object(payload); + const rawModels = object(root?.models); + if (!rawModels) return []; + + const models: ModelDefinition[] = []; + const seen = new Set(); + + for (const [modelId, raw] of Object.entries(rawModels)) { + if (ANTIGRAVITY_DENYLIST.has(modelId)) continue; + const model = object(raw); + if (!model) continue; + + const id = modelId.trim(); + if (!id || seen.has(id)) continue; + seen.add(id); + + const displayName = text(model.displayName) ?? id; + const supportsThinking = model.supportsThinking === true; + const reasoningEfforts = supportsThinking ? ["low", "medium", "high"] : []; + const maxOutputTokens = typeof model.maxOutputTokens === "number" && model.maxOutputTokens > 0 + ? model.maxOutputTokens + : 65_536; + + models.push({ + id, + displayName, + capabilities: { + images: model.supportsImages === true || id.includes("gemini") || id.includes("claude"), + }, + maxOutputTokens, + privateData: { reasoningEfforts }, + }); + } + + // Merge static models from Antigravity catalog that might not be dynamically returned + for (const staticModel of STATIC_ANTIGRAVITY_MODELS) { + if (!seen.has(staticModel.id)) { + seen.add(staticModel.id); + models.push(staticModel); + } + } + + return models; +} + +export function reasoningEfforts(model: ModelSnapshot): string[] { + const data = object(model.privateData); + const efforts = data?.reasoningEfforts; + return Array.isArray(efforts) ? efforts.filter((item) => typeof item === "string") : []; +} + +export const antigravityModels: ModelSupport = { + list: async ({ resource }, context): Promise => { + if (!resource) return STATIC_ANTIGRAVITY_MODELS; + let data; + try { + data = accountData(resource); + } catch { + return STATIC_ANTIGRAVITY_MODELS; + } + + const payloads = [ + JSON.stringify({ project: data.projectId || "bamboo-precept-lgxtn" }), + JSON.stringify({}), + ]; + + 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, + }, + body: bodyPayload, + }, + ); + if (response.status >= 200 && response.status < 300) { + const body = JSON.parse(response.body); + const models = parseAntigravityModels(body); + if (models.length > 0) return models; + } + } catch { + // Continue + } + } + } + + return STATIC_ANTIGRAVITY_MODELS; + }, +}; diff --git a/server/plugins/build-in/antigravity-auth/oauth.ts b/server/plugins/build-in/antigravity-auth/oauth.ts new file mode 100644 index 0000000..e0b6e20 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/oauth.ts @@ -0,0 +1,157 @@ +import type { JsonValue, PluginContext } from "cursor-byok:plugin"; +import type { + OAuth2AuthorizationCodeAddMethod, + OAuth2AuthorizationCodeBegin, + ResourceDraft, +} from "cursor-byok:resource"; +import { credentialDraft, queryAccountQuota } from "./resources.ts"; +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"; + +const AUTHORIZATION_LIFETIME_MS = 5 * 60 * 1000; + +type Session = { createdAtMs: number }; + +function object(value: unknown): Record | null { + return value !== null && typeof value === "object" && !Array.isArray(value) + ? value as Record + : null; +} + +function text(value: unknown): string | null { + return typeof value === "string" && value.trim() ? value.trim() : null; +} + +function parseBody(body: string): Record { + try { + return object(JSON.parse(body)) ?? {}; + } catch { + return {}; + } +} + +function parseSession(value: JsonValue): Session { + const session = object(value); + 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 }; +} + +async function begin( + input: { redirectUri: string; state: string; codeChallenge: string }, + _context: PluginContext, +): Promise { + const authParams = new URLSearchParams({ + client_id: CLIENT_ID, + response_type: "code", + redirect_uri: input.redirectUri, + scope: SCOPES.join(" "), + state: input.state, + code_challenge: input.codeChallenge, + code_challenge_method: "S256", + access_type: "offline", + prompt: "consent", + }); + return { + session: { createdAtMs: Date.now() }, + authorizationUrl: `${GOOGLE_AUTHORIZATION_URL}?${authParams.toString()}`, + expiresAtMs: Date.now() + AUTHORIZATION_LIFETIME_MS, + }; +} + +async function complete( + sessionValue: JsonValue, + input: { code: string; redirectUri: string; codeVerifier: string }, + context: PluginContext, +): Promise { + 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}`); + } + + 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 userInfoResponse = await context.network.fetch( + "https://www.googleapis.com/oauth2/v1/userinfo?alt=json", + { + method: "GET", + headers: { authorization: `Bearer ${accessToken}`, accept: "application/json" }, + }, + ); + if (userInfoResponse.status >= 200 && userInfoResponse.status < 300) { + displayName = text(parseBody(userInfoResponse.body).email) ?? displayName; + } + } catch { + // Account identity has a token fingerprint fallback. + } + + 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 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 and Claude models.", + "zh-CN": "使用 Google 账号完成 Antigravity 授权,以使用 Gemini 与 Claude 模型。", + }, + callback: { port: CALLBACK_PORT, path: CALLBACK_PATH }, + begin, + complete, +}; diff --git a/server/plugins/build-in/antigravity-auth/oauth_test.ts b/server/plugins/build-in/antigravity-auth/oauth_test.ts new file mode 100644 index 0000000..5be3ff6 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/oauth_test.ts @@ -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", + ); +}); diff --git a/server/plugins/build-in/antigravity-auth/plugin.json b/server/plugins/build-in/antigravity-auth/plugin.json new file mode 100644 index 0000000..32354f0 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/plugin.json @@ -0,0 +1,19 @@ +{ + "apiVersion": 1, + "id": "dev.cursorbyok.plugins.antigravity-auth", + "name": "Antigravity", + "version": "0.3.0", + "author": "Antigravity", + "minAppVersion": "0.1.0", + "icon": "assets/antigravity.svg", + "entry": "main.ts", + "permissions": { + "network": [ + "daily-cloudcode-pa.googleapis.com", + "daily-cloudcode-pa.sandbox.googleapis.com", + "cloudcode-pa.googleapis.com", + "oauth2.googleapis.com", + "www.googleapis.com" + ] + } +} diff --git a/server/plugins/build-in/antigravity-auth/provider.ts b/server/plugins/build-in/antigravity-auth/provider.ts new file mode 100644 index 0000000..c211e09 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/provider.ts @@ -0,0 +1,703 @@ +import type { + LlmContentPart, + LlmMessage, + LlmRequest, + ProviderInvokeInput, + ProviderOutput, + ProviderResult, + ProviderSupport, +} from "cursor-byok:provider"; +import type { JsonValue, PluginContext } from "cursor-byok:plugin"; +import { HttpError } from "cursor-byok:protocol/openai-chat"; +import { + ANTIGRAVITY_CLIENT_HEADERS, + ANTIGRAVITY_ENDPOINTS, + ANTIGRAVITY_USER_AGENT, + antigravityModels, +} from "./models.ts"; +import { + type AccountData, + accountData, + isTokenExpired, + quotaExhaustedPatch, + refreshAccount, + RESOURCE_TYPE, +} from "./resources.ts"; + +export function isQuotaError(error: string): boolean { + const message = error.toLowerCase(); + return message.includes("resource_exhausted") || + message.includes("quota_exceeded") || + message.includes("quota_exhausted") || + message.includes("rate_limit_exceeded") || + message.includes("rate limit") || + message.includes("model_capacity_exhausted") || + message.includes("too many requests") || + message.includes("429"); +} + +function isQuotaHttpError(error: HttpError): boolean { + if (error.status === 429) return true; + const body = error.body.toLowerCase(); + return body.includes("resource_exhausted") || + body.includes("quota_exceeded") || + body.includes("quota_exhausted") || + body.includes("rate_limit_exceeded") || + body.includes("rate limit") || + body.includes("model_capacity_exhausted") || + body.includes("user rate limit exceeded") || + body.includes("too many requests"); +} + +function invalidResult(message: string, stateMessage: string): ProviderResult { + return { + status: "resource-error", + message, + patch: { state: { status: "invalid", message: stateMessage } }, + }; +} + +async function readBody(lines: AsyncIterable): Promise { + const collected: string[] = []; + for await (const line of lines) collected.push(line); + return collected.join("\n"); +} + +function resolveAntigravityModel(modelId: string): string { + const raw = modelId.trim(); + const lower = raw.toLowerCase(); + + // 1. If explicit tier is already specified in the model ID, pass it directly! + if ( + lower.startsWith("gemini-3.7-flash-") || + lower.startsWith("gemini-3.6-flash-") || + lower.startsWith("gemini-3.1-pro-") || + lower === "gemini-3.7-flash" || + lower === "gemini-3.6-flash" || + lower === "gemini-2.5-flash" || + lower === "gemini-2.5-pro" || + lower === "gemini-2.0-flash" || + lower === "claude-sonnet-4-6" || + lower === "claude-sonnet-4-6-thinking" || + lower === "claude-opus-4-6-thinking" || + lower === "gemini-3.1-flash-image" || + lower === "gpt-oss-120b-medium" + ) { + if (lower === "gemini-3.1-pro-high") return "gemini-pro-agent"; + return raw; + } + + // 2. Canonical Antigravity-Manager mapping for aliases + 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" + ) { + return "claude-opus-4-6-thinking"; + } + 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") { + return "gemini-2.5-flash"; + } + if (lower === "gemini-3-flash" || lower === "gemini-3.5-flash") { + return "gemini-3.7-flash"; + } + if (lower === "gemini-3-pro" || lower === "gemini-3.1-pro") { + return "gemini-3.1-pro-preview"; + } + if (lower === "gemini-3-pro-high") { + return "gemini-pro-agent"; + } + + return raw; +} + +function randomHex(length = 8): string { + const array = new Uint8Array(Math.ceil(length / 2)); + crypto.getRandomValues(array); + return Array.from(array, (byte) => byte.toString(16).padStart(2, "0")).join("").slice(0, length); +} + +function generateRequestId(): string { + return `agent/${Date.now()}/${randomHex(8)}`; +} + +// Keys the CloudCode v1internal Schema proto rejects with "Cannot find field" +const UNSUPPORTED_SCHEMA_KEYS: Record = { + "$schema": true, + "$ref": true, + "$defs": true, + "$comment": true, + "examples": true, + "unevaluatedProperties": true, + "unevaluatedItems": true, + "patternProperties": true, + "propertyNames": true, + "exclusiveMinimum": true, + "exclusiveMaximum": true, + "multipleOf": true, + "dependencies": true, + "dependentSchemas": true, + "dependentRequired": true, + "deprecated": true, + "readOnly": true, + "writeOnly": true, + "x-mcp-header": true, + "const": true, + "default": true, + "additionalProperties": true, + "title": true, + "format": true, +}; + +const PROTO_TYPE_MAP: Record = { + string: "STRING", + number: "NUMBER", + integer: "INTEGER", + boolean: "BOOLEAN", + array: "ARRAY", + object: "OBJECT", +}; + +function enforceUppercaseTypes(value: unknown): unknown { + if (Array.isArray(value)) return value.map(enforceUppercaseTypes); + if (value === null || typeof value !== "object") return value; + const out: Record = {}; + for (const [key, child] of Object.entries(value as Record)) { + if (UNSUPPORTED_SCHEMA_KEYS[key]) continue; + if (key === "type" && typeof child === "string") { + out[key] = PROTO_TYPE_MAP[child.toLowerCase()] ?? child.toUpperCase(); + } else { + out[key] = enforceUppercaseTypes(child); + } + } + if (!out.type && out.properties) { + out.type = "OBJECT"; + } + return out; +} + +function sanitizeSchema(value: unknown): unknown { + if (!value || typeof value !== "object") { + return { type: "OBJECT", properties: {} }; + } + const clean = enforceUppercaseTypes(value) as Record; + if (!clean.type) clean.type = "OBJECT"; + return clean; +} + +function convertToCloudCodeContents( + instructions: string, + messages: LlmMessage[], +): { + contents: Array<{ role: string; parts: Array> }>; + systemInstruction?: { role: string; parts: Array<{ text: string }> }; +} { + const rawContents: Array<{ role: string; parts: Array> }> = []; + let systemText = instructions || ""; + + for (const msg of messages) { + if (msg.role === "system") { + const txt = msg.content + .map((p) => (p.type === "text" ? p.text : "")) + .filter(Boolean) + .join("\n"); + if (txt) { + systemText += (systemText ? "\n\n" : "") + txt; + } + continue; + } + + if (msg.role === "assistant") { + const parts: Array> = []; + if (msg.text) { + parts.push({ text: msg.text }); + } + + const replayVal = msg.replayState?.providerKind === "antigravity" + ? (msg.replayState.value as Record | null) + : 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 + : {}, + }, + thoughtSignature: sig || "skip_thought_signature_validator", + }); + } + + if (parts.length > 0) { + rawContents.push({ role: "model", parts }); + } + } else if (msg.role === "tool") { + rawContents.push({ + role: "user", + parts: [ + { + functionResponse: { + name: msg.name || "function", + response: { result: msg.content }, + }, + }, + ], + }); + } else if (msg.role === "user") { + const parts: Array> = []; + for (const p of msg.content) { + if (p.type === "text") { + if (p.text) parts.push({ text: p.text }); + } else if (p.type === "image") { + parts.push({ + inlineData: { + mimeType: p.mediaType, + data: p.dataBase64, + }, + }); + } + } + if (parts.length === 0) { + parts.push({ text: " " }); + } + rawContents.push({ role: "user", parts }); + } + } + + // Merge consecutive same-role messages so contents strictly alternate user -> model -> user -> model + const contents: Array<{ role: string; parts: Array> }> = []; + for (const item of rawContents) { + if (item.parts.length === 0) continue; + const last = contents[contents.length - 1]; + if (last && last.role === item.role) { + last.parts.push(...item.parts); + } else { + contents.push(item); + } + } + + if (contents.length > 0 && contents[0].role !== "user") { + contents.unshift({ role: "user", parts: [{ text: " " }] }); + } + + return { + contents, + ...(systemText.trim() + ? { systemInstruction: { role: "system", parts: [{ text: systemText.trim() }] } } + : {}), + }; +} + +async function streamCloudCode( + accessToken: string, + projectId: string, + modelId: string, + input: ProviderInvokeInput, + output: ProviderOutput, + context: PluginContext, +): Promise { + const actualModel = resolveAntigravityModel(modelId); + const { contents, systemInstruction } = convertToCloudCodeContents( + input.request.instructions, + input.request.messages, + ); + + 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), + })), + }, + ] + : undefined; + + const toolConfig = tools + ? { + functionCallingConfig: { mode: "AUTO" }, + } + : undefined; + + const payload = { + project: projectId || "bamboo-precept-lgxtn", + model: actualModel, + userAgent: "antigravity", + requestType: "agent", + requestId: generateRequestId(), + enabledCreditTypes: ["GOOGLE_ONE_AI"], + request: { + contents, + ...(systemInstruction ? { systemInstruction } : {}), + ...(tools ? { tools } : {}), + ...(toolConfig ? { toolConfig } : {}), + generationConfig: { + maxOutputTokens: 65536, + }, + }, + }; + + const headers: Record = { + authorization: `Bearer ${accessToken}`, + "content-type": "application/json", + "user-agent": ANTIGRAVITY_USER_AGENT, + ...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"; + } + + let lastError: Error | null = null; + let hasEmittedAnyChunk = false; + + for (const endpoint of ANTIGRAVITY_ENDPOINTS) { + if (hasEmittedAnyChunk) break; + + try { + const response = await context.network.stream( + `${endpoint}/v1internal:streamGenerateContent?alt=sse`, + { + method: "POST", + headers, + body: JSON.stringify(payload), + }, + ); + + 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 + ) { + continue; + } + throw lastError; + } + + let textStarted = false; + let thinkingStarted = false; + let doneEmitted = false; + let hasTools = false; + let toolIndex = 0; + let lastThoughtSignature: string | 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; + const raw = line.slice(5).trim(); + if (!raw || raw === "[DONE]") break; + + let json: Record; + try { + json = JSON.parse(raw) as Record; + } catch { + continue; + } + + const resp = (json.response as Record | undefined) ?? json; + if (!resp) continue; + + const usage = resp.usageMetadata as Record | undefined; + if (usage) { + finalUsage = { + inputTokens: typeof usage.promptTokenCount === "number" ? usage.promptTokenCount : null, + outputTokens: typeof usage.candidatesTokenCount === "number" + ? usage.candidatesTokenCount + : null, + totalTokens: typeof usage.totalTokenCount === "number" ? usage.totalTokenCount : null, + }; + } + + const candidates = resp.candidates as Array> | undefined; + const candidate = candidates?.[0]; + const content = candidate?.content as Record | undefined; + const parts = content?.parts as Array> | undefined; + + if (parts) { + for (const part of parts) { + const sig = typeof part.thoughtSignature === "string" ? part.thoughtSignature : null; + if (sig) { + lastThoughtSignature = sig; + } + + const isThought = part.thought === true; + const textPart = typeof part.text === "string" ? part.text : null; + + if (isThought && textPart) { + const cleanThought = textPart.replace(/<\/?think>/gi, ""); + if (cleanThought) { + hasEmittedAnyChunk = true; + if (!thinkingStarted) { + thinkingStarted = true; + output.emit({ type: "thinking-start" }); + } + output.emit({ type: "thinking-delta", text: cleanThought }); + } + } else if (textPart) { + if (thinkingStarted) { + thinkingStarted = false; + output.emit({ type: "thinking-end" }); + } + const cleanText = textPart.replace(/<\/?think>/gi, ""); + if (cleanText) { + hasEmittedAnyChunk = true; + if (!textStarted) { + textStarted = true; + output.emit({ type: "text-start" }); + } + output.emit({ type: "text-delta", text: cleanText }); + } + } + + const fnCall = part.functionCall as { name: string; args: unknown } | undefined; + if (fnCall) { + hasEmittedAnyChunk = true; + if (thinkingStarted) { + thinkingStarted = false; + output.emit({ type: "thinking-end" }); + } + if (textStarted) { + textStarted = false; + output.emit({ type: "text-end" }); + } + hasTools = true; + const currentIdx = toolIndex++; + const callId = `call_${Date.now()}_${currentIdx}`; + output.emit({ + type: "tool-call-start", + index: currentIdx, + callId, + name: fnCall.name, + }); + const argsStr = typeof fnCall.args === "string" + ? fnCall.args + : JSON.stringify(fnCall.args || {}); + output.emit({ + type: "tool-call-arguments-delta", + index: currentIdx, + delta: argsStr, + }); + output.emit({ + type: "tool-call-end", + index: currentIdx, + }); + } + } + } + + const finishReason = typeof candidate?.finishReason === "string" + ? candidate.finishReason + : null; + if (finishReason) { + if (thinkingStarted) { + thinkingStarted = false; + output.emit({ type: "thinking-end" }); + } + if (textStarted) { + textStarted = false; + output.emit({ type: "text-end" }); + } + if (lastThoughtSignature) { + output.emit({ + type: "replay-state", + providerKind: "antigravity", + value: { thoughtSignature: lastThoughtSignature }, + }); + lastThoughtSignature = null; + } + if (finalUsage) { + output.emit({ + type: "usage", + usage: { + inputTokens: finalUsage.inputTokens, + outputTokens: finalUsage.outputTokens, + totalTokens: finalUsage.totalTokens, + cacheReadTokens: null, + cacheWriteTokens: null, + reasoningTokens: null, + }, + }); + finalUsage = null; + } + const isTool = finishReason === "STOP" && + (hasTools || parts?.some((p) => p.functionCall)); + output.emit({ + type: "done", + reason: isTool ? "tool-use" : "stop", + }); + doneEmitted = true; + break; + } + } + + if (thinkingStarted) { + output.emit({ type: "thinking-end" }); + } + if (textStarted) { + output.emit({ type: "text-end" }); + } + if (lastThoughtSignature) { + output.emit({ + type: "replay-state", + providerKind: "antigravity", + value: { thoughtSignature: lastThoughtSignature }, + }); + } + if (finalUsage) { + output.emit({ + type: "usage", + usage: { + inputTokens: finalUsage.inputTokens, + outputTokens: finalUsage.outputTokens, + totalTokens: finalUsage.totalTokens, + cacheReadTokens: null, + cacheWriteTokens: null, + reasoningTokens: null, + }, + }); + } + if (!doneEmitted) { + output.emit({ + type: "done", + reason: hasTools ? "tool-use" : "stop", + }); + } + return; + } catch (err) { + lastError = err instanceof Error ? err : new Error(String(err)); + if (hasEmittedAnyChunk) { + throw lastError; + } + } + } + + if (lastError) throw lastError; +} + +async function invoke( + input: ProviderInvokeInput, + output: ProviderOutput, + context: PluginContext, +): Promise { + if (!input.resource) { + return { + status: "request-error", + message: "Add a Google Antigravity account or API key before calling Antigravity", + }; + } + let data: AccountData; + try { + data = accountData(input.resource); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return invalidResult(message, message); + } + + let patchData: AccountData | null = null; + + // Auto-refresh token if expired or close to expiration (skew 5 mins) + if (data.refreshToken && isTokenExpired(data)) { + try { + const refreshed = await refreshAccount(input.resource, context); + if (refreshed.privateData) { + data = refreshed.privateData as unknown as AccountData; + patchData = data; + } + } catch { + // Continue with existing token + } + } + + let projectId = data.projectId ?? "bamboo-precept-lgxtn"; + + 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" }; + } catch (error) { + if (error instanceof HttpError) { + 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, + ); + return { + status: "completed", + patch: { privateData: freshData as unknown as JsonValue, state: { status: "ready" } }, + }; + } + } catch { + // Failed refresh + } + } + if (isQuotaHttpError(error)) { + return { + status: "resource-error", + message: error.message, + patch: quotaExhaustedPatch(data, error.body), + }; + } + return { status: "request-error", message: error.message }; + } + const message = error instanceof Error ? error.message : String(error); + if (isQuotaError(message)) { + return { status: "resource-error", message, patch: quotaExhaustedPatch(data, message) }; + } + return { status: "request-error", message }; + } +} + +export const antigravityProvider: ProviderSupport = { + id: "antigravity", + displayName: { + "en-US": "Google Antigravity", + "zh-CN": "Google Antigravity", + }, + description: { + "en-US": "Google Antigravity / Gemini model access with hybrid reasoning & agent tools.", + "zh-CN": "通过 Google Antigravity / Gemini API 使用混合推理与 Agent 工具。", + }, + providerType: "google", + resourceType: RESOURCE_TYPE, + models: antigravityModels, + invoke, +}; diff --git a/server/plugins/build-in/antigravity-auth/resources.ts b/server/plugins/build-in/antigravity-auth/resources.ts new file mode 100644 index 0000000..e463b08 --- /dev/null +++ b/server/plugins/build-in/antigravity-auth/resources.ts @@ -0,0 +1,630 @@ +import type { JsonValue, NetworkRequestInit, PluginContext } from "cursor-byok:plugin"; +import type { + ResourceDraft, + ResourceImportFile, + ResourceImportResult, + ResourceImportSupport, + ResourceMetric, + ResourcePatch, + ResourceSnapshot, + ResourceState, + ResourceView, +} from "cursor-byok:resource"; +import { + ANTIGRAVITY_CLIENT_HEADERS, + 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; +}; + +export type AccountQuota = { + planLabel: string | null; + limitReached: boolean; + coolingUntilMs: number | null; + updatedAtMs: number; + claude?: QuotaMetric | null; + gemini?: QuotaMetric | null; +}; + +export type AccountData = { + accessToken: string; + refreshToken: string | null; + displayName: string; + projectId?: string | null; + expiresAtMs?: number | null; + quota: AccountQuota | null; +}; + +export type CredentialCandidate = { + accessToken: string; + refreshToken: string | null; + displayName: string | null; + projectId?: string | null; + expiresAtMs?: number | null; + quota?: AccountQuota | null; +}; + +export async function fetchAccountProjectAndTier( + accessToken: string, + network: PluginContext["network"], +): Promise<{ projectId: string; planLabel: string }> { + for (const endpoint of ANTIGRAVITY_ENDPOINTS) { + try { + const assistRes = await network.fetch(`${endpoint}/v1internal:loadCodeAssist`, { + method: "POST", + headers: { + authorization: `Bearer ${accessToken}`, + "content-type": "application/json", + "user-agent": ANTIGRAVITY_USER_AGENT, + ...ANTIGRAVITY_CLIENT_HEADERS, + }, + body: JSON.stringify({ metadata: { ideType: "ANTIGRAVITY" } }), + }); + if (assistRes.status >= 200 && assistRes.status < 300) { + const body = object(JSON.parse(assistRes.body)); + 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); + 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"; + } + return { projectId: project ?? "bamboo-precept-lgxtn", planLabel }; + } + } catch { + // Continue next endpoint + } + } + return { projectId: "bamboo-precept-lgxtn", planLabel: "FREE" }; +} + +export async function queryAccountQuota( + accessToken: string, + network: PluginContext["network"], +): Promise<{ quota: AccountQuota | null; projectId: string }> { + const { projectId, planLabel } = await fetchAccountProjectAndTier(accessToken, network); + + for (const endpoint of ANTIGRAVITY_ENDPOINTS) { + try { + const response = await network.fetch(`${endpoint}/v1internal:fetchAvailableModels`, { + method: "POST", + headers: { + authorization: `Bearer ${accessToken}`, + "content-type": "application/json", + accept: "application/json", + "user-agent": ANTIGRAVITY_USER_AGENT, + ...ANTIGRAVITY_CLIENT_HEADERS, + }, + body: JSON.stringify({ project: projectId }), + }); + if (response.status < 200 || response.status >= 300) continue; + const root = object(JSON.parse(response.body)); + const models = object(root?.models); + if (!models) continue; + + let claudeFraction: number | null = null; + let claudeResetAtMs: number | null = null; + let geminiFraction: number | null = null; + let geminiResetAtMs: number | null = null; + + 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 resetTime = text(quota?.resetTime); + const resetAtMs = resetTime ? Date.parse(resetTime) : null; + if (fraction === null) continue; + + const k = key.toLowerCase(); + if (k.includes("claude") || k.includes("sonnet") || k.includes("opus")) { + if (claudeFraction === null || fraction < claudeFraction) { + claudeFraction = fraction; + claudeResetAtMs = resetAtMs; + } + } else if (k.includes("gemini") || k.includes("flash") || k.includes("pro")) { + if (geminiFraction === null || fraction < geminiFraction) { + geminiFraction = fraction; + geminiResetAtMs = resetAtMs; + } + } + } + + return { + projectId, + quota: { + planLabel, + 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, + }, + }; + } catch { + // Continue next endpoint + } + } + + return { + projectId, + quota: { + planLabel, + limitReached: false, + coolingUntilMs: null, + updatedAtMs: Date.now(), + claude: null, + gemini: null, + }, + }; +} + +function object(value: unknown): Record | null { + return value !== null && typeof value === "object" && !Array.isArray(value) + ? value as Record + : null; +} + +function text(value: unknown): string | null { + return typeof value === "string" && value.trim() ? value.trim() : null; +} + +function decodeJwtPayload(token: string): Record | null { + const parts = token.split("."); + if (parts.length < 2) return null; + try { + const normalized = parts[1].replace(/-/g, "+").replace(/_/g, "/"); + const padded = normalized.padEnd(Math.ceil(normalized.length / 4) * 4, "="); + const bytes = Uint8Array.from(atob(padded), (char) => char.charCodeAt(0)); + return object(JSON.parse(new TextDecoder().decode(bytes))); + } catch { + return null; + } +} + +function claim(payload: Record | null, key: string): string | null { + return payload ? text(payload[key]) : null; +} + +export function isJwtExpired(token: string, bufferSeconds = 300): boolean { + if (token.startsWith("AIza") || !token.includes(".")) return false; + const payload = decodeJwtPayload(token); + if (!payload) return false; + const exp = typeof payload.exp === "number" ? payload.exp : null; + if (!exp) return false; + const nowSeconds = Math.floor(Date.now() / 1000); + return exp <= (nowSeconds + bufferSeconds); +} + +export function isTokenExpired(data: AccountData, bufferSeconds = 300): boolean { + if (!data.refreshToken) return false; + if (typeof data.expiresAtMs === "number" && data.expiresAtMs > 0) { + return Date.now() >= data.expiresAtMs - bufferSeconds * 1000; + } + return isJwtExpired(data.accessToken, bufferSeconds); +} + +async function tokenFingerprint(token: string): Promise { + const digest = await crypto.subtle.digest("SHA-256", new TextEncoder().encode(token)); + return Array.from( + new Uint8Array(digest).slice(0, 8), + (byte) => byte.toString(16).padStart(2, "0"), + ).join(""); +} + +export async function accountIdentity( + token: string, + providedDisplayName?: string | null, +): Promise<{ key: string; displayName: string }> { + const payload = decodeJwtPayload(token); + const email = claim(payload, "email"); + const sub = claim(payload, "sub"); + const name = claim(payload, "name") ?? claim(payload, "preferred_username"); + + const fingerprint = await tokenFingerprint(token); + 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); + return { key: `antigravity:${identity}`, displayName }; +} + +export async function credentialDraft(credential: CredentialCandidate): Promise { + const identity = await accountIdentity(credential.accessToken, credential.displayName); + const data: AccountData = { + accessToken: credential.accessToken, + refreshToken: credential.refreshToken, + displayName: credential.displayName ?? identity.displayName, + projectId: credential.projectId ?? "bamboo-precept-lgxtn", + expiresAtMs: credential.expiresAtMs ?? + (credential.refreshToken ? Date.now() + 3500 * 1000 : null), + quota: credential.quota ?? null, + }; + return { key: identity.key, privateData: data as unknown as JsonValue }; +} + +export function accountData(resource: ResourceSnapshot): AccountData { + const data = object(resource.privateData); + const accessToken = text(data?.accessToken); + if (!accessToken) throw new Error("Antigravity account resource is missing its access token"); + return { + accessToken, + refreshToken: text(data?.refreshToken), + displayName: text(data?.displayName) ?? "Antigravity account", + projectId: text(data?.projectId) ?? "bamboo-precept-lgxtn", + expiresAtMs: typeof data?.expiresAtMs === "number" ? data.expiresAtMs : null, + quota: (data?.quota ?? null) as AccountQuota | null, + }; +} + +export function accountHeaders(data: AccountData): Record { + return { + authorization: `Bearer ${data.accessToken}`, + accept: "application/json", + "user-agent": ANTIGRAVITY_USER_AGENT, + ...ANTIGRAVITY_CLIENT_HEADERS, + }; +} + +export function quotaState(quota: AccountQuota | null, nowMs = Date.now()): ResourceState { + if (!quota || !quota.limitReached) return { status: "ready" }; + const coolingUntil = quota.coolingUntilMs; + if (coolingUntil !== null && coolingUntil > nowMs) { + return { + status: "cooling", + retryAtMs: coolingUntil, + message: "Antigravity rate limit reached; cooling down", + }; + } + return { status: "ready" }; +} + +export function quotaExhaustedPatch( + data: AccountData, + error?: string, + nowMs = Date.now(), +): ResourcePatch { + let retryAfterMs = 60 * 1000; + if (error) { + const match = error.match(/retry(?:_after|\s+after)?\s*[:=]?\s*(\d+)/i); + if (match?.[1]) { + const parsed = Number(match[1]); + if (Number.isFinite(parsed) && parsed > 0) { + retryAfterMs = parsed > 10_000_000 ? parsed - nowMs : parsed * 1000; + } + } + } + const coolingUntilMs = nowMs + Math.max(5000, retryAfterMs); + const quota: AccountQuota = { + planLabel: data.quota?.planLabel ?? "Antigravity / Gemini", + limitReached: true, + coolingUntilMs, + updatedAtMs: nowMs, + claude: data.quota?.claude ?? null, + gemini: data.quota?.gemini ?? null, + }; + return { + privateData: { ...data, quota } as unknown as JsonValue, + state: quotaState(quota, nowMs), + }; +} + +export function presentAccount(resource: ResourceSnapshot): ResourceView { + const data = accountData(resource); + const metrics: ResourceMetric[] = []; + if (data.quota?.claude) { + metrics.push({ + id: "claude", + label: { "en-US": "Claude", "zh-CN": "Claude" }, + unit: "percent", + value: data.quota.claude.remainingPercent, + ...(data.quota.claude.resetAtMs ? { resetAtMs: data.quota.claude.resetAtMs } : {}), + }); + } + if (data.quota?.gemini) { + metrics.push({ + id: "gemini", + label: { "en-US": "Gemini", "zh-CN": "Gemini" }, + unit: "percent", + value: data.quota.gemini.remainingPercent, + ...(data.quota.gemini.resetAtMs ? { resetAtMs: data.quota.gemini.resetAtMs } : {}), + }); + } + return { + displayName: data.displayName, + ...(data.quota?.planLabel ? { description: data.quota.planLabel } : {}), + ...(metrics.length > 0 ? { metrics } : {}), + }; +} + +export async function refreshAccount( + resource: ResourceSnapshot, + context: PluginContext, +): Promise { + const data = accountData(resource); + let accessToken = data.accessToken; + let refreshToken = data.refreshToken; + let projectId = data.projectId ?? "bamboo-precept-lgxtn"; + let expiresAtMs = data.expiresAtMs ?? null; + + if (refreshToken) { + const response = await context.network.fetch(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, + grant_type: "refresh_token", + refresh_token: refreshToken, + }).toString(), + }); + + if (response.status < 200 || response.status >= 300) { + const bodyText = response.body.toLowerCase(); + // 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", + }, + }; + } + // On network glitches or temporary Google server errors, keep ready + return { + state: { status: "ready" }, + }; + } + + const body = object(JSON.parse(response.body)); + accessToken = text(body?.access_token) ?? accessToken; + refreshToken = text(body?.refresh_token) ?? refreshToken; + const expiresIn = typeof body?.expires_in === "number" ? body.expires_in : 3600; + expiresAtMs = Date.now() + expiresIn * 1000; + } + + // Fetch real-time quota and project ID + const result = await queryAccountQuota(accessToken, context.network); + projectId = result.projectId || projectId; + + const updatedData: AccountData = { + ...data, + accessToken, + refreshToken, + projectId, + expiresAtMs, + quota: result.quota ?? data.quota, + }; + return { + privateData: updatedData as unknown as JsonValue, + state: { status: "ready" }, + }; +} + +function firstText(source: Record, keys: string[]): string | null { + for (const key of keys) { + const value = text(source[key]); + if (value) return value; + } + return null; +} + +function collectCredentials(value: unknown, output: CredentialCandidate[]): void { + if (Array.isArray(value)) { + for (const item of value) collectCredentials(item, output); + return; + } + const item = object(value); + if (!item || item.disabled === true) return; + for (const key of ["accounts", "credentials", "items", "keys"]) { + if (Array.isArray(item[key])) { + collectCredentials(item[key], output); + return; + } + } + const tokens = object(item.tokens) ?? item; + let accessToken = firstText(tokens, [ + "access", + "accessToken", + "access_token", + "token", + "apiKey", + "api_key", + "key", + "GEMINI_API_KEY", + "GOOGLE_API_KEY", + "ANTIGRAVITY_API_KEY", + ]) ?? firstText(item, [ + "access", + "accessToken", + "access_token", + "token", + "apiKey", + "api_key", + "key", + "GEMINI_API_KEY", + "GOOGLE_API_KEY", + "ANTIGRAVITY_API_KEY", + ]); + const refreshToken = firstText(tokens, ["refresh", "refresh_token", "refreshToken"]) ?? + 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"]); + + if (!accessToken && !refreshToken) return; + if (!accessToken && refreshToken) { + accessToken = refreshToken; + } + if (!accessToken) return; + output.push({ accessToken, refreshToken, displayName, projectId }); +} + +export async function parseCredentialFiles( + files: ResourceImportFile[], + network?: PluginContext["network"], +): Promise<{ + credentials: CredentialCandidate[]; + warnings: string[]; +}> { + const credentials: CredentialCandidate[] = []; + const warnings: string[] = []; + for (const file of files) { + const raw = file.content.trim(); + if (!raw) continue; + + // 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 }); + continue; + } + + // Try parsing as JSON + let content: unknown; + 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, + ); + if (envMatch?.[1]) { + credentials.push({ + accessToken: envMatch[1].trim(), + refreshToken: null, + displayName: file.name, + }); + continue; + } + const keyMatch = raw.match(/AIza[0-9A-Za-z-_]{35}/); + if (keyMatch?.[0]) { + credentials.push({ accessToken: keyMatch[0], refreshToken: null, displayName: file.name }); + continue; + } + warnings.push(`${file.name}: not valid JSON or API key`); + continue; + } + + if (typeof content === "string") { + credentials.push({ accessToken: content.trim(), refreshToken: null, displayName: file.name }); + continue; + } + + const found: CredentialCandidate[] = []; + collectCredentials(content, found); + if (found.length === 0) { + warnings.push(`${file.name}: no Google/Antigravity API key or token found`); + continue; + } + for (const candidate of found) { + if (candidate.refreshToken && candidate.accessToken === candidate.refreshToken) { + try { + const response = await fetchText(network, 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, + grant_type: "refresh_token", + refresh_token: candidate.refreshToken, + }).toString(), + }); + const body = object(JSON.parse(response.body)); + 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); + } + } catch { + // Keep placeholder + } + } + credentials.push(candidate); + } + } + return { credentials, warnings }; +} + +export const credentialImport: ResourceImportSupport = { + displayName: { + "en-US": "Import Google / Antigravity Credentials", + "zh-CN": "导入 Google / Antigravity 凭证", + }, + description: { + "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 => { + 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 or API key", + ); + } + const drafts = await Promise.all( + credentials.map(async (c) => { + try { + const res = await queryAccountQuota(c.accessToken, context.network); + c.quota = res.quota; + c.projectId = res.projectId; + } catch { + // ignore error + } + return credentialDraft(c); + }), + ); + return { + resources: drafts, + ...(warnings.length > 0 ? { warnings } : {}), + }; + }, +}; diff --git a/server/src/plugin/builtin.rs b/server/src/plugin/builtin.rs index 9ddd5ff..f2ca588 100644 --- a/server/src/plugin/builtin.rs +++ b/server/src/plugin/builtin.rs @@ -109,7 +109,70 @@ const GROK_AUTH: &[(&str, &str)] = &[ ), ]; -const PLUGINS: &[(&str, &[(&str, &str)])] = &[("codex-auth", CODEX_AUTH), ("grok-auth", GROK_AUTH)]; +const ANTIGRAVITY_AUTH: &[(&str, &str)] = &[ + ( + "plugin.json", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/plugin.json" + )), + ), + ( + "main.ts", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/main.ts" + )), + ), + ( + "provider.ts", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/provider.ts" + )), + ), + ( + "models.ts", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/models.ts" + )), + ), + ( + "oauth.ts", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/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!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/resources.ts" + )), + ), + ( + "assets/antigravity.svg", + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/plugins/build-in/antigravity-auth/assets/antigravity.svg" + )), + ), +]; + +const PLUGINS: &[(&str, &[(&str, &str)])] = &[ + ("codex-auth", CODEX_AUTH), + ("grok-auth", GROK_AUTH), + ("antigravity-auth", ANTIGRAVITY_AUTH), +]; /// 把内置插件预装到 installed 目录。manifest 的 version 是缓存键: /// 版本一致时零写盘;版本变化时整目录同步并清理旧版本残留文件。 @@ -209,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(); diff --git a/server/src/plugin/catalog.rs b/server/src/plugin/catalog.rs index 69fda85..d9fddb7 100644 --- a/server/src/plugin/catalog.rs +++ b/server/src/plugin/catalog.rs @@ -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(()) diff --git a/server/src/plugin/descriptor.rs b/server/src/plugin/descriptor.rs index 2ac33dd..2762255 100644 --- a/server/src/plugin/descriptor.rs +++ b/server/src/plugin/descriptor.rs @@ -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, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub struct OAuthCallbackDefinition { + #[serde(default)] + pub port: Option, + #[serde(default)] + pub path: Option, } #[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)] diff --git a/server/src/plugin/mod.rs b/server/src/plugin/mod.rs index 7789c40..48ccac6 100644 --- a/server/src/plugin/mod.rs +++ b/server/src/plugin/mod.rs @@ -7,6 +7,7 @@ mod definition; mod descriptor; mod installation; mod manifest; +mod oauth_callback; mod protocol; mod registry; mod runtime; diff --git a/server/src/plugin/oauth_callback.rs b/server/src/plugin/oauth_callback.rs new file mode 100644 index 0000000..69661ae --- /dev/null +++ b/server/src/plugin/oauth_callback.rs @@ -0,0 +1,347 @@ +//! 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, + Router, +}; +use serde::Deserialize; +use tokio::sync::{oneshot, Mutex}; +use tokio_util::sync::CancellationToken; + +use crate::{Error, Result}; + +const CALLBACK_RESPONSE_TIMEOUT: Duration = Duration::from_secs(120); + +pub(super) struct CallbackRequest { + pub result: std::result::Result, + pub response: oneshot::Sender, +} + +#[derive(Debug)] +pub(super) struct CallbackOutcome { + pub success: bool, + pub message: Option, +} + +pub(super) struct CallbackHandle { + pub redirect_uri: String, + pub receiver: oneshot::Receiver, + shutdown: CancellationToken, +} + +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>>>, + shutdown: CancellationToken, +} + +#[derive(Deserialize)] +struct CallbackQuery { + code: Option, + state: Option, + error: Option, + error_description: Option, +} + +pub(super) async fn bind( + port: Option, + path: &str, + expected_state: String, + plugin_name: String, + plugin_icon: String, + resource_name: serde_json::Value, +) -> Result { + 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, + headers: HeaderMap, + Query(query): Query, +) -> Html { + 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.", + )), + )), + } +} + +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#" + + +{title}

{title}

{plugin_name}
{resource_name}
{detail}
{close}
"# + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn escapes_plugin_content_in_callback_page() { + let state = CallbackState { + expected_state: "state".into(), + plugin_name: "".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("")); + } + + #[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")); + } +} diff --git a/server/src/plugin/registry.rs b/server/src/plugin/registry.rs index 29e22bc..c81cb29 100644 --- a/server/src/plugin/registry.rs +++ b/server/src/plugin/registry.rs @@ -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,8 +14,9 @@ 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, @@ -53,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, pub verification_url: String, pub verification_url_complete: Option, pub expires_at_ms: i64, @@ -288,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 { - 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 { @@ -339,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 { @@ -348,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 { let executable = self.executable()?; let entry = self.find_entry(&executable, &plugin_id).await?; let value = self @@ -386,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); @@ -413,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 { + 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 = 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, + ) -> Result { + 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, @@ -899,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, +} + +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 { diff --git a/server/src/plugin/sdk/plugin.ts b/server/src/plugin/sdk/plugin.ts index 8bb9ab0..4106784 100644 --- a/server/src/plugin/sdk/plugin.ts +++ b/server/src/plugin/sdk/plugin.ts @@ -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 ? { diff --git a/server/src/plugin/sdk/resource.ts b/server/src/plugin/sdk/resource.ts index 2322909..7c15acc 100644 --- a/server/src/plugin/sdk/resource.ts +++ b/server/src/plugin/sdk/resource.ts @@ -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; + complete( + session: JsonValue, + input: { + code: string; + redirectUri: string; + codeVerifier: string; + }, + context: PluginContext, + ): Promise; +}; + +export type OAuth2AuthorizationCodeBegin = { + session: JsonValue; + authorizationUrl: string; + expiresAtMs: number; + pollIntervalMs?: number; +}; + +export type ResourceAddMethod = OAuth2AddMethod | OAuth2AuthorizationCodeAddMethod; export type ResourceImportFile = { name: string; diff --git a/server/src/plugin/sdk/worker.ts b/server/src/plugin/sdk/worker.ts index 2e29368..a08849b 100644 --- a/server/src/plugin/sdk/worker.ts +++ b/server/src/plugin/sdk/worker.ts @@ -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": {