fix: adjust secondary button height and enhance MultiCombobox input behavior

This commit is contained in:
leokun
2026-08-24 18:06:49 +08:00
parent fe17e15e75
commit 89631bdab0
22 changed files with 754 additions and 84 deletions
+1
View File
@@ -8,6 +8,7 @@ export interface Provider {
name: string;
provider_type: ProviderType;
base_url: string;
api_key?: string;
has_api_key: boolean;
custom_headers: Record<string, string | null>;
extra_params: Record<string, unknown>;
@@ -1,5 +1,5 @@
import type { ProviderInput, ProviderType } from "../api";
import { FormField, TextInput } from "./ui/FormControls";
import { FormField, SecretTextInput, TextInput } from "./ui/FormControls";
import { JsonEditor } from "./ui/JsonEditor";
import { Select } from "./ui/Select";
import { claudeIcon, openAiIcon } from "./ui/icons";
@@ -23,7 +23,7 @@ export function ProviderEditor({ value, headersText, extraText, editing, onChang
{ value: "anthropic", label: "Anthropic", icon: claudeIcon },
]} onChange={(provider_type) => patch({ provider_type: provider_type as ProviderType })} /></FormField>
<FormField label="Base URL" hint={t("模型服务的 API 根地址;修改后会同步更新该上游模型的路由身份。")}><TextInput placeholder="https://api.example.com/v1" value={value.base_url} onChange={(event) => patch({ base_url: event.target.value })} /></FormField>
<FormField className={styles.fullWidth} label="API Key" hint={editing ? t("留空表示保留当前 API Key。") : t("访问模型服务所需的密钥。")}><TextInput type="password" autoComplete="off" placeholder={editing ? t("留空以保留当前密钥") : "sk-xxxxxx"} value={value.api_key ?? ""} onChange={(event) => patch({ api_key: event.target.value })} /></FormField>
<FormField className={styles.fullWidth} label="API Key" hint={editing ? t("留空表示保留当前 API Key。") : t("访问模型服务所需的密钥。")}><SecretTextInput autoComplete="off" placeholder={editing ? t("留空以保留当前密钥") : "sk-xxxxxx"} value={value.api_key ?? ""} onChange={(event) => patch({ api_key: event.target.value })} /></FormField>
<FormField className={styles.fullWidth} label={t("自定义 Headers JSON")} hint={t("值必须是字符串;编辑时 null 表示保留对应敏感 Header 的原值。")}><JsonEditor ariaLabel={t("自定义 Headers JSON")} value={headersText} onChange={onHeadersChange} /></FormField>
<FormField className={styles.fullWidth} label={t("额外参数 JSON")} hint={t("合并到该上游所有模型的请求体。")}><JsonEditor ariaLabel={t("额外参数 JSON")} value={extraText} onChange={onExtraChange} /></FormField>
</div>;
@@ -1,5 +1,5 @@
import type { ModelInput, Provider, ProviderInput, ProviderType } from "../../api";
import { FormField, TextInput } from "../ui/FormControls";
import { FormField, SecretTextInput, TextInput } from "../ui/FormControls";
import { Checkbox } from "../ui/Checkbox";
import { JsonEditor } from "../ui/JsonEditor";
import { Combobox, MultiCombobox, Select } from "../ui/Select";
@@ -72,7 +72,7 @@ export function CursorModelEditor({ draft, providers, editing, modelOptions, dis
<div className={styles.grid}>
{!editing && draft.providerMode === "new" && <>
<FormField label="Base URL" hint={t("模型服务的 API 根地址,例如 https://api.openai.com/v1。")}><TextInput placeholder="例如:https://api.openai.com/v1" value={draft.provider.base_url} onChange={(event) => setProvider({ base_url: event.target.value })} /></FormField>
<FormField label="API Key" hint={t("访问模型服务所需的密钥。")}><TextInput type="password" placeholder="例如:sk-xxxxxx" autoComplete="off" value={draft.provider.api_key ?? ""} onChange={(event) => setProvider({ api_key: event.target.value })} /></FormField>
<FormField label="API Key" hint={t("访问模型服务所需的密钥。")}><SecretTextInput placeholder="例如:sk-xxxxxx" autoComplete="off" value={draft.provider.api_key ?? ""} onChange={(event) => setProvider({ api_key: event.target.value })} /></FormField>
</>}
<FormField label={t("端点类型")} hint={t("默认继承上游,可为当前模型单独修改。")}><Select ariaLabel={t("端点类型")} value={draft.model.endpoint_type} options={[
{ value: "openai-responses", label: "OpenAI Responses", icon: openAiIcon }, { value: "openai-chat", label: "OpenAI Chat", icon: openAiIcon }, { value: "anthropic", label: "Anthropic", icon: claudeIcon },
@@ -11,8 +11,8 @@
}
.secondary {
min-height: 34px;
height: 34px;
min-height: 30px;
height: 30px;
display: inline-flex;
align-items: center;
gap: 6px;
@@ -48,3 +48,31 @@
height: 34px;
padding: 0 10px;
}
.secret {
position: relative;
width: 100%;
input {
padding-right: 34px;
}
}
.secretToggle {
position: absolute;
top: 0;
right: 0;
width: 30px;
height: 34px;
display: grid;
place-items: center;
padding: 0;
color: var(--vscode-descriptionForeground);
background: transparent;
border: 0;
cursor: pointer;
&:hover {
color: var(--vscode-foreground);
}
}
@@ -1,13 +1,23 @@
import type { InputHTMLAttributes } from "react";
import { useState, type InputHTMLAttributes } from "react";
import { Icon } from "./Icon";
import { TooltipTrigger } from "./TooltipTrigger";
import { informationOutlineIcon } from "./icons";
import { eyeIcon, eyeOffIcon, informationOutlineIcon } from "./icons";
import styles from "./FormControls.module.scss";
export function TextInput(props: InputHTMLAttributes<HTMLInputElement>) {
return <input {...props} className={[styles.input, props.className].filter(Boolean).join(" ")} />;
}
export function SecretTextInput({ className, ...props }: InputHTMLAttributes<HTMLInputElement>) {
const [visible, setVisible] = useState(false);
return <div className={styles.secret}>
<input {...props} type={visible ? "text" : "password"} className={[styles.input, className].filter(Boolean).join(" ")} />
<button type="button" className={styles.secretToggle} aria-label={visible ? t("隐藏 API Key") : t("显示 API Key")} onClick={() => setVisible((current) => !current)}>
<Icon icon={visible ? eyeOffIcon : eyeIcon} size="1.1em" />
</button>
</div>;
}
export function FormField({ label, hint, className, children }: { label: string; hint?: string; className?: string; children: React.ReactNode }) {
return <label className={[styles.field, className].filter(Boolean).join(" ")}>
<div className={styles.label}>
+1 -1
View File
@@ -178,7 +178,7 @@ export function MultiCombobox({ value, options = [], placeholder, disabled, appe
return <div className={styles.comboRow}><div ref={root} className={styles.multiCombo} data-open={open || undefined}>
<div className={styles.multiValues}>
{value.length > 0 && <span className={styles.multiCount}>{t("已选择 {count} 个", { count: value.length })}</span>}
<input ref={input} value={query} placeholder={value.length ? t("继续选择或输入") : placeholder} disabled={disabled} role="combobox" aria-haspopup="listbox" aria-controls={open ? menuId : undefined} aria-expanded={open} aria-autocomplete="list" onFocus={() => { if (options.length) setOpen(true); }} onChange={(event) => { setQuery(event.target.value); setActive(0); if (options.length) setOpen(true); }} onKeyDown={(event) => {
<input ref={input} value={query} placeholder={value.length ? t("继续选择或输入") : placeholder} disabled={disabled} role="combobox" aria-haspopup="listbox" aria-controls={open ? menuId : undefined} aria-expanded={open} aria-autocomplete="list" onFocus={() => { if (options.length) setOpen(true); }} onBlur={() => add(query)} onChange={(event) => { setQuery(event.target.value); setActive(0); if (options.length) setOpen(true); }} onKeyDown={(event) => {
if (event.key === "ArrowDown") { event.preventDefault(); move(1); }
if (event.key === "ArrowUp") { event.preventDefault(); move(-1); }
if (event.key === "Enter") {
+1
View File
@@ -27,6 +27,7 @@ export const settingsIcon = icon('<path fill="currentColor" fill-rule="evenodd"
export const addIcon = icon('<path fill="currentColor" d="M19 13h-6v6h-2v-6H5v-2h6V5h2v6h6z"/>'); // mdi:plus
export const editIcon = icon('<path fill="currentColor" d="m14.06 9l.94.94L5.92 19H5v-.92zm3.6-6c-.25 0-.51.1-.7.29l-1.83 1.83l3.75 3.75l1.83-1.83c.39-.39.39-1.04 0-1.41l-2.34-2.34c-.2-.2-.45-.29-.71-.29m-3.6 3.19L3 17.25V21h3.75L17.81 9.94z"/>'); // mdi:pencil-outline
export const eyeIcon = icon('<path fill="currentColor" d="M12 9a3 3 0 0 1 3 3a3 3 0 0 1-3 3a3 3 0 0 1-3-3a3 3 0 0 1 3-3m0-4.5c5 0 9.27 3.11 11 7.5c-1.73 4.39-6 7.5-11 7.5S2.73 16.39 1 12c1.73-4.39 6-7.5 11-7.5M3.18 12a9.821 9.821 0 0 0 17.64 0a9.821 9.821 0 0 0-17.64 0"/>'); // mdi:eye-outline
export const eyeOffIcon = icon('<path fill="currentColor" d="M2 5.27L3.28 4L20 20.72L18.73 22l-3.08-3.08c-1.15.38-2.37.58-3.65.58c-5 0-9.27-3.11-11-7.5c.69-1.76 1.79-3.31 3.19-4.54zM12 9a3 3 0 0 1 3 3a3 3 0 0 1-.17 1L11 9.17A3 3 0 0 1 12 9m0-4.5c5 0 9.27 3.11 11 7.5a11.8 11.8 0 0 1-4 5.19l-1.42-1.43A9.86 9.86 0 0 0 20.82 12A9.82 9.82 0 0 0 12 6.5c-1.09 0-2.16.18-3.16.5L7.3 5.47c1.44-.62 3.03-.97 4.7-.97M3.18 12A9.82 9.82 0 0 0 12 17.5c.69 0 1.37-.07 2-.21L11.72 15A3.064 3.064 0 0 1 9 12.28L5.6 8.87c-.99.85-1.82 1.91-2.42 3.13"/>'); // mdi:eye-off-outline
export const informationOutlineIcon = icon('<path fill="currentColor" d="M11 9h2V7h-2m1 13c-4.41 0-8-3.59-8-8s3.59-8 8-8s8 3.59 8 8s-3.59 8-8 8m0-18A10 10 0 0 0 2 12a10 10 0 0 0 10 10a10 10 0 0 0 10-10A10 10 0 0 0 12 2m-1 15h2v-6h-2z"/>'); // mdi:information-outline
export const cilBadgeIcon = icon('<path fill="currentColor" d="m328.375 384l3.698 74.999l-75.862-52.719l-76.287 52.769L183.625 384h-32.039l-5.522 112h36.692l73.413-50.78L329.242 496h36.694l-5.522-112zm87.034-229.086l-2.194-48.054L372.7 80.933l-25.932-40.519l-48.055-2.2L256 16.093l-42.713 22.126l-48.055 2.2L139.3 80.933L98.785 106.86l-2.194 48.054l-22.127 42.714l22.127 42.715l2.2 48.053l40.509 25.927l25.928 40.52l48.055 2.195L256 379.164l42.713-22.126l48.055-2.195l25.928-40.52l40.518-25.923l2.195-48.053l22.127-42.715Zm-31.646 76.949L382 270.377l-32.475 20.78l-20.78 32.475l-38.515 1.76L256 343.125l-34.234-17.733l-38.515-1.76l-20.78-32.475L130 270.377l-1.759-38.514l-17.741-34.235l17.737-34.228L130 124.88l32.471-20.78l20.78-32.474l38.515-1.76L256 52.132l34.234 17.733l38.515 1.76l20.78 32.474L382 124.88l1.759 38.515l17.741 34.233Z"/>', 512, 512); // cil:badge
export const refreshIcon = icon('<path fill="currentColor" d="M17.65 6.35A7.96 7.96 0 0 0 12 4a8 8 0 0 0-8 8a8 8 0 0 0 8 8c3.73 0 6.84-2.55 7.73-6h-2.08A5.99 5.99 0 0 1 12 18a6 6 0 0 1-6-6a6 6 0 0 1 6-6c1.66 0 3.14.69 4.22 1.78L13 11h7V4z"/>'); // mdi:refresh
+26 -2
View File
@@ -362,7 +362,7 @@
{
"file": "pages/CursorSettingsPage.tsx",
"line": 199,
"column": 205
"column": 182
},
{
"file": "pages/CursorSettingsPage.tsx",
@@ -420,7 +420,7 @@
{
"file": "pages/CursorSettingsPage.tsx",
"line": 199,
"column": 217
"column": 194
}
]
},
@@ -2260,6 +2260,18 @@
}
]
},
"86b7355ec3bd55ef": {
"source": "隐藏 API Key",
"kind": "text",
"placeholders": [],
"refs": [
{
"file": "components/ui/FormControls.tsx",
"line": 15,
"column": 81
}
]
},
"8716e1344b0daddb": {
"source": "Cursor 官方",
"kind": "text",
@@ -4030,6 +4042,18 @@
}
]
},
"f2bdc88464c51c2e": {
"source": "显示 API Key",
"kind": "text",
"placeholders": [],
"refs": [
{
"file": "components/ui/FormControls.tsx",
"line": 15,
"column": 99
}
]
},
"f396118b8afd2a21": {
"source": "Cursor 接管已生效;添加上游及其模型配置后即可使用 BYOK 模型。",
"kind": "text",
+2
View File
@@ -156,6 +156,7 @@
"83fcfb4c1f2c1641": "Fetch models",
"842b9f11cdd96bda": "Launch at login",
"84924374710e03bd": "Base URL must be a valid URL",
"86b7355ec3bd55ef": "Hide API Key",
"8716e1344b0daddb": "Cursor official",
"878a8ab176429a86": "View instructions",
"883cc47637fe70f3": "Custom request headers appended to every request for this provider. Values must be strings.",
@@ -281,6 +282,7 @@
"ee239f3943293f87": "Sunday",
"ee6b89a6a740a4c4": "If a port is occupied, a new random port is selected and saved automatically. Restart the app after changing these settings.",
"f04c91a6bc3a6926": "Extra parameters",
"f2bdc88464c51c2e": "Show API Key",
"f396118b8afd2a21": "Cursor interception is active. Add a provider and its model configuration to use BYOK models.",
"f3a76d896853c1df": "Miss",
"f4694c46b1e19602": "Final request type",
+2
View File
@@ -156,6 +156,7 @@
"83fcfb4c1f2c1641": "获取模型",
"842b9f11cdd96bda": "开机启动",
"84924374710e03bd": "Base URL 必须是有效地址",
"86b7355ec3bd55ef": "隐藏 API Key",
"8716e1344b0daddb": "Cursor 官方",
"878a8ab176429a86": "查看说明",
"883cc47637fe70f3": "附加到该上游所有请求的自定义请求头,值必须是字符串。",
@@ -281,6 +282,7 @@
"ee239f3943293f87": "周日",
"ee6b89a6a740a4c4": "端口被占用时会自动选择新的随机端口并保存。修改后需要重启软件才会生效。",
"f04c91a6bc3a6926": "额外参数",
"f2bdc88464c51c2e": "显示 API Key",
"f396118b8afd2a21": "Cursor 接管已生效;添加上游及其模型配置后即可使用 BYOK 模型。",
"f3a76d896853c1df": "未命中",
"f4694c46b1e19602": "最终请求类型",
+1 -1
View File
@@ -40,7 +40,7 @@ export function ProvidersPage() {
};
const openEdit = (provider: Provider) => {
setEditing(provider);
setDraft({ name: provider.name, provider_type: provider.provider_type, base_url: provider.base_url, api_key: "", custom_headers: provider.custom_headers, extra_params: provider.extra_params });
setDraft({ name: provider.name, provider_type: provider.provider_type, base_url: provider.base_url, api_key: provider.api_key ?? "", custom_headers: provider.custom_headers, extra_params: provider.extra_params });
setHeadersText(JSON.stringify(provider.custom_headers, null, 2));
setExtraText(JSON.stringify(provider.extra_params, null, 2));
};
+154 -49
View File
@@ -18,7 +18,7 @@ use crate::{
ContentPart, CursorRunTraceArtifact, CursorRunTraceSummary, LlmCallRequest,
LlmCallResponseChunk, LlmCallSummary, ModelInvocation, ModelRequest, ModelSpec, Overview,
ProjectedContent, ProjectedMessage, PromptSpec, ProviderEndpoint, ProviderEndpointInput,
ProviderEndpointSecret, ProviderModel, ProviderModelInput, ProviderType, Role,
ProviderModel, ProviderModelInput, ProviderType, Role,
},
provider::{ModelEvent, Provider},
store::{
@@ -349,34 +349,15 @@ impl ControlService {
pub async fn discover_input(&self, input: &ProviderEndpointInput) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let endpoint = ProviderEndpoint {
provider_id: 0,
name: input.name.clone(),
provider_type: input.provider_type,
base_url: crate::model::normalize_base_url(&input.base_url)?,
has_api_key: input
.api_key
.as_deref()
.is_some_and(|value| !value.is_empty()),
custom_headers: input.custom_headers.clone(),
extra_params: input.extra_params.clone(),
created_at_ms: 0,
updated_at_ms: 0,
};
let secret = ProviderEndpointSecret {
endpoint,
api_key: input.api_key.clone().unwrap_or_default(),
custom_headers: input.custom_headers.clone(),
};
let mut models = match input.provider_type {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
openai_models(&client, &secret).await?
}
ProviderType::Anthropic => anthropic_models(&client, &secret).await?,
};
models.sort();
models.dedup();
Ok(DiscoveredModels { models })
let base_url = crate::model::normalize_base_url(&input.base_url)?;
discover_provider_models(
&client,
input.provider_type,
&base_url,
input.api_key.as_deref().unwrap_or_default(),
&input.custom_headers,
)
.await
}
pub async fn discover_models(&self, provider_id: i64) -> Result<DiscoveredModels> {
@@ -386,15 +367,14 @@ impl ControlService {
.provider(provider_id)
.await?
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
let mut models = match provider.endpoint.provider_type {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
openai_models(&client, &provider).await?
}
ProviderType::Anthropic => anthropic_models(&client, &provider).await?,
};
models.sort();
models.dedup();
Ok(DiscoveredModels { models })
discover_provider_models(
&client,
provider.endpoint.provider_type,
&provider.endpoint.base_url,
provider.endpoint.api_key.as_deref().unwrap_or_default(),
&provider.custom_headers,
)
.await
}
pub async fn calls(&self, limit: i64) -> Result<Vec<CallSummary>> {
@@ -606,15 +586,51 @@ fn readable_utf8(data: &[u8]) -> Option<&str> {
.then_some(value)
}
async fn discover_provider_models(
client: &reqwest::Client,
provider_type: ProviderType,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<DiscoveredModels> {
let mut models = match provider_type {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
openai_models(client, base_url, api_key, custom_headers).await?
}
ProviderType::Anthropic => {
anthropic_models(client, base_url, api_key, custom_headers).await?
}
};
models.sort();
models.dedup();
Ok(DiscoveredModels { models })
}
fn model_discovery_url(base_url: &str) -> Result<Url> {
let mut url = Url::parse(base_url)
.map_err(|error| Error::Config(format!("invalid provider base URL: {error}")))?;
if url.host_str().is_none() {
return Err(Error::Config(
"provider base URL must contain a host".into(),
));
}
url.set_path("/v1/models");
url.set_query(None);
url.set_fragment(None);
Ok(url)
}
async fn openai_models(
client: &reqwest::Client,
provider: &ProviderEndpointSecret,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut request = client.get(format!("{}/models", provider.endpoint.base_url));
if !provider.api_key.is_empty() {
request = request.bearer_auth(&provider.api_key);
let mut request = client.get(model_discovery_url(base_url)?);
if !api_key.is_empty() {
request = request.bearer_auth(api_key);
}
let response = apply_custom_headers(request, &provider.custom_headers)?
let response = apply_discovery_headers(request, custom_headers)?
.send()
.await?;
let status = response.status();
@@ -629,22 +645,24 @@ async fn openai_models(
async fn anthropic_models(
client: &reqwest::Client,
provider: &ProviderEndpointSecret,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut after_id = None::<String>;
let mut found = BTreeSet::new();
loop {
let mut request = client
.get(format!("{}/models", provider.endpoint.base_url))
.get(model_discovery_url(base_url)?)
.query(&[("limit", "100")])
.header("anthropic-version", "2023-06-01");
if !provider.api_key.is_empty() {
request = request.header("x-api-key", &provider.api_key);
if !api_key.is_empty() {
request = request.header("x-api-key", api_key);
}
if let Some(after_id) = &after_id {
request = request.query(&[("after_id", after_id)]);
}
let response = apply_custom_headers(request, &provider.custom_headers)?
let response = apply_discovery_headers(request, custom_headers)?
.send()
.await?;
let status = response.status();
@@ -699,7 +717,7 @@ fn estimate_output_tokens(output: &str) -> u64 {
}
}
fn apply_custom_headers(
fn apply_discovery_headers(
mut request: reqwest::RequestBuilder,
headers: &serde_json::Value,
) -> Result<reqwest::RequestBuilder> {
@@ -707,6 +725,9 @@ fn apply_custom_headers(
.as_object()
.ok_or_else(|| Error::Config("custom headers must be an object".into()))?;
for (name, value) in object {
if name.eq_ignore_ascii_case("user-agent") {
continue;
}
let value = value
.as_str()
.ok_or_else(|| Error::Config(format!("custom header {name} must be a string")))?;
@@ -838,4 +859,88 @@ mod tests {
assert_eq!(super::estimate_output_tokens("1 2 3"), 3);
assert_eq!(super::estimate_output_tokens(""), 0);
}
#[test]
fn model_discovery_url_uses_only_the_provider_origin() {
assert_eq!(
super::model_discovery_url("https://example.com:8443/arbitrary/v1/chat/completions")
.unwrap()
.as_str(),
"https://example.com:8443/v1/models"
);
}
#[tokio::test]
async fn model_discovery_does_not_inherit_user_agent_or_request_body_settings() {
type CapturedRequest = (
axum::http::Method,
axum::http::Uri,
axum::http::HeaderMap,
bytes::Bytes,
);
async fn models(
axum::extract::State(sender): axum::extract::State<
tokio::sync::mpsc::UnboundedSender<CapturedRequest>,
>,
request: axum::extract::Request,
) -> axum::Json<serde_json::Value> {
let (parts, body) = request.into_parts();
let body = axum::body::to_bytes(body, usize::MAX).await.unwrap();
sender
.send((parts.method, parts.uri, parts.headers, body))
.unwrap();
axum::Json(serde_json::json!({ "data": [{ "id": "model-a" }] }))
}
let (sender, mut requests) = tokio::sync::mpsc::unbounded_channel();
let app = axum::Router::new()
.route("/v1/models", axum::routing::get(models))
.with_state(sender);
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 directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("discovery.db").display()
))
.await
.unwrap();
let service = ControlService::new(
store,
Arc::new(TestProvider {
invocation: Arc::new(Mutex::new(None)),
}),
)
.unwrap();
let result = service
.discover_input(&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiResponses,
base_url: format!("http://{address}/custom/responses"),
api_key: Some("secret".into()),
custom_headers: serde_json::json!({
"uSeR-aGeNt": "inherited-user-agent",
"x-tenant": "tenant-a"
}),
extra_params: serde_json::json!({ "temperature": 0.7 }),
})
.await
.unwrap();
assert_eq!(result.models, vec!["model-a"]);
let (method, uri, headers, body) = requests.recv().await.unwrap();
assert_eq!(method, axum::http::Method::GET);
assert_eq!(uri.path(), "/v1/models");
assert!(body.is_empty());
assert!(headers.get(axum::http::header::USER_AGENT).is_none());
assert_eq!(headers.get("x-tenant").unwrap(), "tenant-a");
assert_eq!(
headers.get(axum::http::header::AUTHORIZATION).unwrap(),
"Bearer secret"
);
server.abort();
}
}
+147
View File
@@ -459,6 +459,7 @@ pub fn dynamic_mcp(
Error::Protocol(format!("MCP tool {} is missing input schema", wire.name))
})?),
};
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
let name = model_tool_name(&wire.name);
let definition = ToolDefinition {
name: name.clone(),
@@ -477,6 +478,43 @@ pub fn dynamic_mcp(
Ok(output)
}
fn normalize_mcp_parameters(tool_name: &str, mut parameters: Value) -> Result<Value> {
let schema = parameters
.as_object_mut()
.ok_or_else(|| invalid_mcp_parameters(tool_name))?;
match schema.get("type") {
Some(Value::String(schema_type)) if schema_type == "object" => return Ok(parameters),
Some(_) => return Err(invalid_mcp_parameters(tool_name)),
None => {}
}
let object_only_union = ["anyOf", "oneOf"].into_iter().any(|keyword| {
schema
.get(keyword)
.and_then(Value::as_array)
.is_some_and(|branches| {
!branches.is_empty()
&& branches.iter().all(|branch| {
branch
.as_object()
.and_then(|branch| branch.get("type"))
.and_then(Value::as_str)
== Some("object")
})
})
});
if !object_only_union {
return Err(invalid_mcp_parameters(tool_name));
}
schema.insert("type".into(), Value::String("object".into()));
Ok(parameters)
}
fn invalid_mcp_parameters(tool_name: &str) -> Error {
Error::Protocol(format!(
"MCP tool {tool_name} input schema must describe an object"
))
}
fn model_tool_name(name: &str) -> String {
name.chars()
.map(|character| {
@@ -575,6 +613,115 @@ mod tests {
.contains("duplicate MCP tool name after normalization: server_name-tool"));
}
#[test]
fn dynamic_mcp_normalizes_cursor_object_union_without_mutating_wire_schema() {
let original_schema = serde_json::json!({
"$schema": "https://json-schema.org/draft/2020-12/schema",
"anyOf": [
{
"type": "object",
"properties": {
"rootPath": { "type": "string", "minLength": 1 }
},
"required": ["rootPath"],
"additionalProperties": false
},
{
"type": "object",
"properties": {
"rootPaths": {
"type": "array",
"items": { "type": "string", "minLength": 1 },
"minItems": 1
}
},
"required": ["rootPaths"],
"additionalProperties": false
}
]
});
let original_json = original_schema.to_string();
let mut tool = direct_mcp_tool("cursor-app-control-move_agent_to_cloned_root");
tool.input_schema_json = Some(original_json.clone());
let request = pb::AgentRunRequest {
mcp_tools: Some(pb::McpTools {
mcp_tools: vec![tool],
}),
..Default::default()
};
let tools = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap();
let (wire, definition) = tools
.get("cursor-app-control-move_agent_to_cloned_root")
.unwrap();
assert_eq!(definition.parameters["type"], "object");
assert_eq!(definition.parameters["anyOf"], original_schema["anyOf"]);
assert_eq!(
wire.input_schema_json.as_deref(),
Some(original_json.as_str())
);
}
#[test]
fn dynamic_mcp_preserves_valid_object_schema() {
let original_schema = serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" }
},
"required": ["query"],
"additionalProperties": false
});
let mut tool = direct_mcp_tool("search");
tool.input_schema_json = Some(original_schema.to_string());
let request = pb::AgentRunRequest {
mcp_tools: Some(pb::McpTools {
mcp_tools: vec![tool],
}),
..Default::default()
};
let tools = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap();
let (_, definition) = tools.get("search").unwrap();
assert_eq!(definition.parameters, original_schema);
}
#[test]
fn dynamic_mcp_rejects_schemas_that_are_not_provably_objects() {
let invalid_schemas = [
serde_json::Value::Null,
serde_json::json!({ "type": "string" }),
serde_json::json!({ "properties": { "query": { "type": "string" } } }),
serde_json::json!({
"anyOf": [
{ "type": "object" },
{ "type": "string" }
]
}),
];
for schema in invalid_schemas {
let mut tool = direct_mcp_tool("unsafe_schema");
tool.input_schema_json = Some(schema.to_string());
let request = pb::AgentRunRequest {
mcp_tools: Some(pb::McpTools {
mcp_tools: vec![tool],
}),
..Default::default()
};
let error = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap_err();
assert!(
error
.to_string()
.contains("MCP tool unsafe_schema input schema must describe an object"),
"unexpected error for {schema}: {error}"
);
}
}
#[test]
fn meta_mcp_routes_projects_descriptor_routing_without_runtime_discovery() {
let context = pb::RequestContext {
+20 -5
View File
@@ -129,11 +129,19 @@ impl CursorSession {
checkpoint_worker_open = false;
}
Input::Completion(completion) => {
self.forward_completion(completion, &mut completions)
.await?;
if let Some(completion) = self
.forward_completion(completion, &mut completions)
.await?
{
ready.push_back(completion);
}
}
Input::CompletionResult(Some(result)) => {
self.forward_completion(result?, &mut completions).await?;
if let Some(completion) =
self.forward_completion(result?, &mut completions).await?
{
ready.push_back(completion);
}
}
Input::CompletionResult(None) => {
return Err(Error::Protocol("tool result channel closed".into()));
@@ -578,7 +586,7 @@ impl CursorSession {
&self,
mut completion: ToolCompletion,
completions: &mut HashMap<String, ToolCompletion>,
) -> Result<()> {
) -> Result<Option<ToolCompletion>> {
if let Some(image) = completion.take_read_image() {
let blob_id = self.store.put_blob(&image.data, &[]).await?;
completion.persist_read_image(&blob_id, &image)?;
@@ -600,7 +608,14 @@ impl CursorSession {
.commands
.send(ClientCommand::ToolResult(result.clone()))
.await
.map_err(|_| Error::RunNotFound(self.context.request_id.clone()))
.map_err(|_| Error::RunNotFound(self.context.request_id.clone()))?;
let Some(dispatched) = self.tools.continue_after(&result.call_id).await? else {
return Ok(None);
};
for message in dispatched.messages {
self.handle.emit(&message)?;
}
Ok(dispatched.completion)
}
async fn forward_injection(&mut self, action: pb::InjectContextAction) -> Result<()> {
+7
View File
@@ -20,6 +20,13 @@ pub(crate) fn path(call: &ToolCall) -> Result<String> {
string(call, field)
}
pub(crate) fn execution_path(call: &ToolCall) -> Result<Option<String>> {
match normalized(&call.name).as_str() {
"write" | "strreplace" | "editnotebook" => path(call).map(Some),
_ => Ok(None),
}
}
pub(crate) fn after_read(
call: &ToolCall,
result: &pb::ReadResult,
+62 -9
View File
@@ -1,11 +1,19 @@
use std::collections::{BTreeMap, HashSet};
use std::{
collections::{BTreeMap, HashSet},
sync::Arc,
};
use tokio::sync::Mutex;
pub mod codec;
mod dispatch;
pub(crate) mod edit;
pub(crate) mod result;
pub mod runtime;
mod schedule;
pub(crate) mod stream;
#[cfg(test)]
mod tests;
use crate::{
model::{CanonicalMessage, MessageContent, Role, ToolCall},
@@ -14,6 +22,7 @@ use crate::{
};
use self::result::{ToolCompletion, ToolResultSender};
use self::schedule::{DeferredEdit, EditSchedule};
use super::{interaction, proto::agent::v1 as pb};
use runtime::{CursorToolRuntime, ExecContext};
@@ -23,6 +32,7 @@ pub struct ToolDispatcher {
results: ToolResultSender,
search: WebSearch,
fetch: WebFetch,
edit_schedule: Arc<Mutex<EditSchedule>>,
}
pub struct DispatchedTool {
@@ -54,6 +64,7 @@ impl ToolDispatcher {
results,
search: WebSearch::built_in(),
fetch: WebFetch::built_in(),
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
}
}
@@ -74,20 +85,62 @@ impl ToolDispatcher {
if state.completed.contains(&call.call_id) {
continue;
}
let message_index = first_tool_index + position;
let publish_started = !state.started.contains(&call.call_id);
let edit_path = if dynamic_mcp.contains_key(&call.name) {
None
} else {
edit::execution_path(call)?
};
if let Some(path) = edit_path {
let next = self.edit_schedule.lock().await.start_or_defer(
path,
DeferredEdit {
call: call.clone(),
message_index,
publish_started,
context: context.clone(),
},
);
let Some(next) = next else {
continue;
};
dispatched.push(
self.start(
&next.call,
next.message_index,
next.publish_started,
dynamic_mcp,
&next.context,
)
.await?,
);
continue;
}
dispatched.push(
self.start(
call,
first_tool_index + position,
!state.started.contains(&call.call_id),
dynamic_mcp,
context,
)
.await?,
self.start(call, message_index, publish_started, dynamic_mcp, context)
.await?,
);
}
Ok(dispatched)
}
pub(crate) async fn continue_after(&self, call_id: &str) -> Result<Option<DispatchedTool>> {
let next = self.edit_schedule.lock().await.complete(call_id)?;
let Some(next) = next else {
return Ok(None);
};
self.start(
&next.call,
next.message_index,
next.publish_started,
&BTreeMap::new(),
&next.context,
)
.await
.map(Some)
}
async fn start(
&self,
call: &ToolCall,
+68
View File
@@ -0,0 +1,68 @@
use std::collections::{HashMap, VecDeque};
use crate::{model::ToolCall, Error, Result};
use super::runtime::ExecContext;
#[derive(Default)]
pub(super) struct EditSchedule {
paths: HashMap<String, EditPathQueue>,
active_paths: HashMap<String, String>,
}
struct EditPathQueue {
active_call_id: String,
waiting: VecDeque<DeferredEdit>,
}
pub(super) struct DeferredEdit {
pub call: ToolCall,
pub message_index: usize,
pub publish_started: bool,
pub context: ExecContext,
}
impl EditSchedule {
pub fn start_or_defer(&mut self, path: String, edit: DeferredEdit) -> Option<DeferredEdit> {
if let Some(queue) = self.paths.get_mut(&path) {
queue.waiting.push_back(edit);
return None;
}
self.active_paths
.insert(edit.call.call_id.clone(), path.clone());
self.paths.insert(
path,
EditPathQueue {
active_call_id: edit.call.call_id.clone(),
waiting: VecDeque::new(),
},
);
Some(edit)
}
pub fn complete(&mut self, call_id: &str) -> Result<Option<DeferredEdit>> {
let Some(path) = self.active_paths.remove(call_id) else {
return Ok(None);
};
let queue = self.paths.get_mut(&path).ok_or_else(|| {
Error::Protocol(format!("active edit path disappeared for call {call_id}"))
})?;
if queue.active_call_id != call_id {
return Err(Error::Protocol(format!(
"edit path is active for {}, not {call_id}",
queue.active_call_id
)));
}
match queue.waiting.pop_front() {
Some(next) => {
queue.active_call_id = next.call.call_id.clone();
self.active_paths.insert(next.call.call_id.clone(), path);
Ok(Some(next))
}
None => {
self.paths.remove(&path);
Ok(None)
}
}
}
}
+139
View File
@@ -0,0 +1,139 @@
use super::*;
use serde_json::json;
fn edit_call(index: usize, call_id: &str, path: &str, old: &str, new: &str) -> ToolCall {
ToolCall {
index,
call_id: call_id.into(),
model_call_id: "model:0".into(),
name: "StrReplace".into(),
arguments_text: String::new(),
arguments: json!({
"path": path,
"old_string": old,
"new_string": new,
}),
}
}
#[tokio::test]
async fn same_path_edits_start_one_at_a_time() {
let runtime = CursorToolRuntime::default();
let dispatcher = ToolDispatcher::new(runtime.clone());
let calls = [
edit_call(0, "first", "/tmp/a.txt", "left", "LEFT"),
edit_call(1, "second", "/tmp/a.txt", "right", "RIGHT"),
edit_call(2, "other", "/tmp/b.txt", "other", "OTHER"),
];
let dispatched = dispatcher
.start_batch(
&calls,
ToolBatchState {
completed: &HashSet::new(),
started: &HashSet::new(),
response_text: "",
response_thinking: "",
},
&[],
&BTreeMap::new(),
&ExecContext::default(),
)
.await
.unwrap();
assert_eq!(dispatched.len(), 2);
assert_eq!(exec(&dispatched[0]).exec_id, "first");
assert_eq!(exec(&dispatched[1]).exec_id, "other");
let mut file = "left right\n".to_string();
let first_write = advance_read(&runtime, exec(&dispatched[0]).id, &file).await;
file = write_text(&first_write);
assert_eq!(file, "LEFT right\n");
complete_write(&runtime, &first_write).await;
let second = dispatcher
.continue_after("first")
.await
.unwrap()
.expect("second same-path edit should start after the first completes");
assert_eq!(exec(&second).exec_id, "second");
let second_write = advance_read(&runtime, exec(&second).id, &file).await;
file = write_text(&second_write);
assert_eq!(file, "LEFT RIGHT\n");
complete_write(&runtime, &second_write).await;
assert!(dispatcher.continue_after("second").await.unwrap().is_none());
}
fn exec(dispatched: &DispatchedTool) -> &pb::ExecServerMessage {
dispatched
.messages
.iter()
.find_map(|message| match message.message.as_ref() {
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => Some(exec),
_ => None,
})
.expect("dispatched edit should contain an Exec request")
}
async fn advance_read(
runtime: &CursorToolRuntime,
id: u32,
content: &str,
) -> pb::ExecServerMessage {
let event = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
output: Some(pb::read_success::Output::Content(content.into())),
..Default::default()
})),
},
)),
..Default::default()
},
runtime,
)
.await
.unwrap();
let codec::ClientExecEvent::Message(message) = event else {
panic!("edit read should advance to a write")
};
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message else {
panic!("edit read should emit an Exec write request")
};
exec
}
fn write_text(exec: &pb::ExecServerMessage) -> String {
let Some(pb::exec_server_message::Message::WriteArgs(args)) = exec.message.as_ref() else {
panic!("expected WriteArgs")
};
args.file_text.clone()
}
async fn complete_write(runtime: &CursorToolRuntime, exec: &pb::ExecServerMessage) {
let Some(pb::exec_server_message::Message::WriteArgs(args)) = exec.message.as_ref() else {
panic!("expected WriteArgs")
};
let event = codec::client_event(
&pb::ExecClientMessage {
id: exec.id,
message: Some(pb::exec_client_message::Message::WriteResult(
pb::WriteResult {
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
path: args.path.clone(),
..Default::default()
})),
},
)),
..Default::default()
},
runtime,
)
.await
.unwrap();
assert!(matches!(event, codec::ClientExecEvent::Completed(_)));
}
+1 -1
View File
@@ -51,6 +51,7 @@ pub struct ProviderEndpoint {
pub name: String,
pub provider_type: ProviderType,
pub base_url: String,
pub api_key: Option<String>,
pub has_api_key: bool,
pub custom_headers: serde_json::Value,
pub extra_params: serde_json::Value,
@@ -61,7 +62,6 @@ pub struct ProviderEndpoint {
#[derive(Clone, Debug)]
pub struct ProviderEndpointSecret {
pub endpoint: ProviderEndpoint,
pub api_key: String,
pub custom_headers: serde_json::Value,
}
+1 -1
View File
@@ -82,7 +82,7 @@ impl Provider for ProviderRouter {
ProviderType::Anthropic => ProviderKind::Anthropic,
},
request_url,
api_key: endpoint.api_key,
api_key: endpoint.endpoint.api_key.clone().unwrap_or_default(),
custom_headers: custom_headers(&endpoint.custom_headers)?,
max_output_tokens: model.max_output_tokens,
request_timeout,
+75 -7
View File
@@ -127,10 +127,15 @@ impl Store {
.provider(provider_id)
.await?
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
let api_key = input.api_key.as_deref().unwrap_or(&current.api_key);
let api_key = input
.api_key
.as_deref()
.or(current.endpoint.api_key.as_deref())
.unwrap_or_default();
let custom_headers = merge_custom_headers(&current.custom_headers, &input.custom_headers)?;
let base_url = normalize_base_url(&input.base_url)?;
let identity_changed = base_url != current.endpoint.base_url || api_key != current.api_key;
let identity_changed = base_url != current.endpoint.base_url
|| api_key != current.endpoint.api_key.as_deref().unwrap_or_default();
let models = if identity_changed {
sqlx::query("SELECT * FROM provider_models WHERE provider_id = ?")
.bind(provider_id)
@@ -276,7 +281,7 @@ impl Store {
for input in inputs {
let hash = model_hash(
&provider.endpoint.base_url,
&provider.api_key,
provider.endpoint.api_key.as_deref().unwrap_or_default(),
input.endpoint_type,
&input.model_id,
)?;
@@ -333,7 +338,7 @@ impl Store {
.expect("model provider must exist");
let next_hash = model_hash(
&provider.endpoint.base_url,
&provider.api_key,
provider.endpoint.api_key.as_deref().unwrap_or_default(),
input.endpoint_type,
&input.model_id,
)?;
@@ -485,6 +490,7 @@ fn validate_model_batch(inputs: &[ProviderModelInput]) -> Result<()> {
fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
let api_key: String = row.try_get("api_key")?;
let has_api_key = !api_key.is_empty();
let headers: serde_json::Value = serde_json::from_str(row.try_get("custom_headers_json")?)?;
let extra_params: serde_json::Value = serde_json::from_str(row.try_get("extra_params_json")?)?;
Ok(ProviderEndpoint {
@@ -492,7 +498,8 @@ fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
name: row.try_get("name")?,
provider_type: ProviderType::from_str(row.try_get("provider_type")?)?,
base_url: row.try_get("base_url")?,
has_api_key: !api_key.is_empty(),
api_key: has_api_key.then_some(api_key),
has_api_key,
custom_headers: redact_custom_headers(&headers),
extra_params,
created_at_ms: row.try_get("created_at_ms")?,
@@ -501,12 +508,10 @@ fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
}
fn secret_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpointSecret> {
let api_key: String = row.try_get("api_key")?;
let custom_headers: serde_json::Value =
serde_json::from_str(row.try_get("custom_headers_json")?)?;
Ok(ProviderEndpointSecret {
endpoint: endpoint_from_row(row)?,
api_key,
custom_headers,
})
}
@@ -847,6 +852,69 @@ mod tests {
assert_eq!(detached, None);
}
#[tokio::test]
async fn updating_provider_without_changing_api_key_preserves_model_hashes() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("provider-key-keep.db").display()
))
.await
.unwrap();
let (created_provider, original) = store
.create_provider_with_model(&provider(), &model("model-a"))
.await
.unwrap();
// Editor keeps the configured key: sending it back must not rehash models.
store
.update_provider(created_provider.provider_id, &provider())
.await
.unwrap();
assert!(store
.provider_model(&original.model_hash)
.await
.unwrap()
.is_some());
// Editor cleared the field: keep the current key, still no rehash.
let mut without_key = provider();
without_key.api_key = None;
store
.update_provider(created_provider.provider_id, &without_key)
.await
.unwrap();
assert!(store
.provider_model(&original.model_hash)
.await
.unwrap()
.is_some());
assert_eq!(store.provider_models(false).await.unwrap().len(), 1);
}
#[tokio::test]
async fn providers_expose_the_configured_api_key_for_editing() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("provider-key-echo.db").display()
))
.await
.unwrap();
let created = store.create_provider(&provider()).await.unwrap();
assert_eq!(created.api_key.as_deref(), Some("secret"));
let listed = store.providers().await.unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].api_key.as_deref(), Some("secret"));
assert!(listed[0].has_api_key);
let without_key = ProviderEndpointInput { api_key: None, ..provider() };
let empty = store.create_provider(&without_key).await.unwrap();
assert_eq!(empty.api_key, None);
assert!(!empty.has_api_key);
assert_eq!(store.providers().await.unwrap().len(), 2);
}
async fn insert_call(store: &Store, provider: &ProviderEndpoint, model: &ProviderModel) {
sqlx::query(
"INSERT INTO llm_calls(call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type, provider_url, request_type, request_url, model_id, display_name, status, created_at_ms, message_count, tool_count, detailed) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",