fix: harden provider, tool, and desktop behavior

This commit is contained in:
leokun
2026-08-24 15:37:42 +08:00
parent 4bb4ca5c8d
commit 7fa4953883
10 changed files with 376 additions and 51 deletions
+9 -2
View File
@@ -1,3 +1,5 @@
LOCAL_TAURI_SIGNING_KEY := $(CURDIR)/.tauri/cursor-byok.local.key
.PHONY: check dev-web dev-server dev-desktop build-web build-server build-desktop build-docker
check:
@@ -21,8 +23,13 @@ build-web:
build-server:
cargo build --release --package cursor-server --bin cursor-server
build-desktop:
npm --prefix apps/desktop run tauri:build
$(LOCAL_TAURI_SIGNING_KEY):
@install -d -m 700 "$(dir $@)"
@apps/desktop/node_modules/.bin/tauri signer generate --ci --write-keys "$@" >/dev/null
@chmod 600 "$@" "$@.pub"
build-desktop: $(LOCAL_TAURI_SIGNING_KEY)
TAURI_SIGNING_PRIVATE_KEY="$(LOCAL_TAURI_SIGNING_KEY)" TAURI_SIGNING_PRIVATE_KEY_PASSWORD="" npm --prefix apps/desktop run tauri:build
build-docker:
docker build --tag cursor-byok:local .
@@ -174,6 +174,7 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
const shouldResetScrollRef = useRef(false)
const scrollApiRef = useRef<ScrollAreaApi | null>(null)
const scrollStateRef = useRef<ScrollAreaState | null>(null)
const contentElementRef = useRef<HTMLDivElement | null>(null)
const spacerRef = useRef<HTMLDivElement | null>(null)
const [contentInsets, setContentInsets] = useState<ContentInsets>({
top: 0,
@@ -192,7 +193,8 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
})
const [, forceUpdate] = useState(0)
const setContentRef = useCallback((node: HTMLDivElement | null) => {
const readContentInsets = useCallback(() => {
const node = contentElementRef.current
const styles = node ? getComputedStyle(node) : null
const nextInsets = {
top: styles ? Number.parseFloat(styles.paddingTop) || 0 : 0,
@@ -205,6 +207,26 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
)
}, [])
const setContentRef = useCallback((node: HTMLDivElement | null) => {
contentElementRef.current = node
readContentInsets()
}, [readContentInsets])
useLayoutEffect(() => {
const node = contentElementRef.current
if (!node) return
readContentInsets()
const resizeObserver = new ResizeObserver(readContentInsets)
resizeObserver.observe(node)
const frame = requestAnimationFrame(readContentInsets)
return () => {
cancelAnimationFrame(frame)
resizeObserver.disconnect()
}
}, [readContentInsets])
const contentInsetTop = contentInsets.top
if (!scrollStateRef.current) {
+10 -1
View File
@@ -7,7 +7,7 @@ use crate::{
Error, Result,
};
use super::{mcp_state, ReadImage, ToolCompletion};
use super::{gate, mcp_state, ReadImage, ToolCompletion};
use crate::cursor::tools::{
edit,
runtime::{ExecStage, PendingExec},
@@ -18,6 +18,15 @@ pub(crate) fn from_exec(
wire_result: &pb::exec_client_message::Message,
) -> Result<ToolCompletion> {
use pb::{exec_client_message::Message, tool_call::Tool};
let mut gated_shell = matches!(
wire_result,
Message::ShellResult(_) | Message::MiniSweAgentBashResult(_)
)
.then(|| wire_result.clone());
if let Some(message) = gated_shell.as_mut() {
gate::exec_message(message);
}
let wire_result = gated_shell.as_ref().unwrap_or(wire_result);
if let Message::McpStateExecResult(result) = wire_result {
return mcp_state::complete(pending, result);
}
+169
View File
@@ -0,0 +1,169 @@
use crate::cursor::proto::agent::v1 as pb;
const KIB: usize = 1024;
const SHELL_STREAM_LIMIT: usize = 16 * KIB;
const SHELL_CONTENT_LIMIT: usize = 32 * KIB;
pub(super) fn model_content(tool: &pb::tool_call::Tool, content: &mut String) {
if matches!(tool, pb::tool_call::Tool::ShellToolCall(_)) {
*content = truncate_edges("Shell", content, SHELL_CONTENT_LIMIT);
}
}
pub(super) fn exec_message(message: &mut pb::exec_client_message::Message) {
use pb::exec_client_message::Message;
match message {
Message::ShellResult(result) | Message::MiniSweAgentBashResult(result) => {
gate_shell_result(result)
}
_ => {}
}
}
fn gate_shell_result(result: &mut pb::ShellResult) {
use pb::shell_result::Result;
match result.result.as_mut() {
Some(Result::Success(success)) => {
success.stdout = truncate_edges("Shell stdout", &success.stdout, SHELL_STREAM_LIMIT);
success.stderr = truncate_edges("Shell stderr", &success.stderr, SHELL_STREAM_LIMIT);
if let Some(interleaved) = success.interleaved_output.as_mut() {
*interleaved =
truncate_edges("Shell interleaved output", interleaved, SHELL_CONTENT_LIMIT);
}
}
Some(Result::Failure(failure)) => {
failure.stdout = truncate_edges("Shell stdout", &failure.stdout, SHELL_STREAM_LIMIT);
failure.stderr = truncate_edges("Shell stderr", &failure.stderr, SHELL_STREAM_LIMIT);
if let Some(interleaved) = failure.interleaved_output.as_mut() {
*interleaved =
truncate_edges("Shell interleaved output", interleaved, SHELL_CONTENT_LIMIT);
}
}
_ => {}
}
}
fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
if content.len() <= limit {
return content.to_string();
}
let original = content.len();
let mut shown = limit;
loop {
let notice = format!(
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; omitted middle; showing {shown} of {original} bytes]\n\n"
);
let available = limit.saturating_sub(notice.len());
let head = utf8_prefix(content, available / 2);
let tail = utf8_suffix(content, available.saturating_sub(head.len()));
let next_shown = head.len().saturating_add(tail.len());
if next_shown == shown {
return format!("{head}{notice}{tail}");
}
shown = next_shown;
}
}
fn utf8_prefix(value: &str, limit: usize) -> &str {
let mut end = limit.min(value.len());
while end > 0 && !value.is_char_boundary(end) {
end -= 1;
}
&value[..end]
}
fn utf8_suffix(value: &str, limit: usize) -> &str {
let mut start = value.len().saturating_sub(limit);
while start < value.len() && !value.is_char_boundary(start) {
start += 1;
}
&value[start..]
}
#[cfg(test)]
mod tests {
use super::*;
fn shell_tool() -> pb::tool_call::Tool {
pb::tool_call::Tool::ShellToolCall(pb::ShellToolCall::default())
}
#[test]
fn shell_output_keeps_both_ends_within_its_budget() {
let mut content = format!("HEAD{}TAIL", " ".repeat(1024 * KIB));
model_content(&shell_tool(), &mut content);
assert!(content.len() <= SHELL_CONTENT_LIMIT);
assert!(content.starts_with("HEAD"));
assert!(content.ends_with("TAIL"));
assert!(content.contains("omitted middle"));
}
#[test]
fn non_shell_output_is_unchanged() {
let mut content = "x".repeat(64 * KIB);
let original = content.clone();
model_content(
&pb::tool_call::Tool::ReadToolCall(pb::ReadToolCall::default()),
&mut content,
);
assert_eq!(content, original);
}
#[test]
fn shell_streams_are_limited_before_rendering() {
let mut message = pb::exec_client_message::Message::ShellResult(pb::ShellResult {
result: Some(pb::shell_result::Result::Success(pb::ShellSuccess {
stdout: format!("HEAD{}TAIL", "x".repeat(64 * KIB)),
stderr: format!("ERROR_HEAD{}ERROR_TAIL", "y".repeat(64 * KIB)),
interleaved_output: Some(format!("START{}END", "z".repeat(64 * KIB))),
..Default::default()
})),
..Default::default()
});
exec_message(&mut message);
let pb::exec_client_message::Message::ShellResult(result) = message else {
panic!("expected Shell result");
};
let Some(pb::shell_result::Result::Success(success)) = result.result else {
panic!("expected Shell success");
};
assert!(success.stdout.len() <= SHELL_STREAM_LIMIT);
assert!(success.stdout.starts_with("HEAD"));
assert!(success.stdout.ends_with("TAIL"));
assert!(success.stderr.len() <= SHELL_STREAM_LIMIT);
assert!(success.stderr.starts_with("ERROR_HEAD"));
assert!(success.stderr.ends_with("ERROR_TAIL"));
assert!(success.interleaved_output.unwrap().len() <= SHELL_CONTENT_LIMIT);
}
#[test]
fn failed_shell_streams_are_limited() {
let mut message = pb::exec_client_message::Message::ShellResult(pb::ShellResult {
result: Some(pb::shell_result::Result::Failure(pb::ShellFailure {
stdout: "x".repeat(64 * KIB),
stderr: "y".repeat(64 * KIB),
interleaved_output: Some("z".repeat(64 * KIB)),
..Default::default()
})),
..Default::default()
});
exec_message(&mut message);
let pb::exec_client_message::Message::ShellResult(result) = message else {
panic!("expected Shell result");
};
let Some(pb::shell_result::Result::Failure(failure)) = result.result else {
panic!("expected Shell failure");
};
assert!(failure.stdout.len() <= SHELL_STREAM_LIMIT);
assert!(failure.stderr.len() <= SHELL_STREAM_LIMIT);
assert!(failure.interleaved_output.unwrap().len() <= SHELL_CONTENT_LIMIT);
}
}
+3 -1
View File
@@ -1,5 +1,6 @@
mod await_shell;
mod exec;
mod gate;
mod interaction;
mod local;
mod mcp;
@@ -87,9 +88,10 @@ impl ToolCompletion {
pub(crate) fn new(
call: &ToolCall,
started_at_ms: u64,
result: ToolResult,
mut result: ToolResult,
tool: pb::tool_call::Tool,
) -> Self {
gate::model_content(&tool, &mut result.content);
Self {
result,
tool_call: pb::ToolCall {
+28 -4
View File
@@ -138,7 +138,12 @@ pub fn normalize_base_url(value: &str) -> Result<String> {
Ok(url.as_str().trim_end_matches('/').to_string())
}
pub fn model_hash(base_url: &str, provider_type: ProviderType, model_id: &str) -> Result<String> {
pub fn model_hash(
base_url: &str,
api_key: &str,
provider_type: ProviderType,
model_id: &str,
) -> Result<String> {
let base_url = normalize_base_url(base_url)?;
let model_id = model_id.trim();
if model_id.is_empty() {
@@ -147,6 +152,8 @@ pub fn model_hash(base_url: &str, provider_type: ProviderType, model_id: &str) -
let mut digest = Sha256::new();
digest.update(base_url.as_bytes());
digest.update([0]);
digest.update(api_key.as_bytes());
digest.update([0]);
digest.update(provider_type.as_str().as_bytes());
digest.update([0]);
digest.update(model_id.as_bytes());
@@ -212,24 +219,41 @@ mod tests {
use super::*;
#[test]
fn hash_uses_normalized_url_type_and_model_only() {
fn hash_uses_normalized_url_key_type_and_model() {
let first = model_hash(
"HTTPS://Example.COM/v1/",
"secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap();
let second = model_hash(
"https://example.com/v1",
"secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap();
assert_eq!(first, second);
assert_eq!(first, "f246010a");
assert_ne!(
first,
model_hash("https://example.com/v1", ProviderType::Anthropic, "model-a").unwrap()
model_hash(
"https://example.com/v1",
"different-secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap()
);
assert_ne!(
first,
model_hash(
"https://example.com/v1",
"secret",
ProviderType::Anthropic,
"model-a",
)
.unwrap()
);
}
-26
View File
@@ -204,32 +204,6 @@ impl Provider for OpenAiResponsesProvider {
}
"response.completed" => {
if let Some(usage) = value.pointer("/response/usage") { yield ModelEvent::Usage(responses_usage(usage)); }
if let Some(output) = value.pointer("/response/output").and_then(Value::as_array) {
for (index, item) in output.iter().enumerate() {
match item.get("type").and_then(Value::as_str) {
Some("reasoning") => {
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
if !reasoning_items.iter().any(|existing| existing.get("id") == item.get("id")) {
reasoning_items.push(item.clone());
}
}
Some("message") => {
if let Some(final_text) = response_item_text(item) {
for event in reconcile_response_text(&mut text_open, &mut text, &final_text) { yield event; }
}
}
Some("function_call") => {
saw_tool = true;
let arguments = item
.get("arguments")
.and_then(Value::as_str)
.map_or(ResponseToolArguments::None, ResponseToolArguments::Snapshot);
for event in update_response_tool(index, item, arguments, true, &mut tools)? { yield event; }
}
_ => {}
}
}
}
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
if text_open { text_open = false; yield ModelEvent::TextEnd; }
for (index, tool) in tools.iter_mut().filter(|(_, tool)| tool.started && !tool.ended) {
+87 -4
View File
@@ -36,7 +36,12 @@ impl Store {
let mut hashes = Vec::with_capacity(models.len());
let mut unique_hashes = HashSet::with_capacity(models.len());
for model in models {
let hash = model_hash(&base_url, model.endpoint_type, &model.model_id)?;
let hash = model_hash(
&base_url,
provider.api_key.as_deref().unwrap_or_default(),
model.endpoint_type,
&model.model_id,
)?;
if !unique_hashes.insert(hash.clone()) {
return Err(Error::Config(format!(
"8-character model hash collision: {hash}"
@@ -125,8 +130,8 @@ impl Store {
let api_key = input.api_key.as_deref().unwrap_or(&current.api_key);
let custom_headers = merge_custom_headers(&current.custom_headers, &input.custom_headers)?;
let base_url = normalize_base_url(&input.base_url)?;
let base_url_changed = base_url != current.endpoint.base_url;
let models = if base_url_changed {
let identity_changed = base_url != current.endpoint.base_url || api_key != current.api_key;
let models = if identity_changed {
sqlx::query("SELECT * FROM provider_models WHERE provider_id = ?")
.bind(provider_id)
.fetch_all(&self.pool)
@@ -140,7 +145,7 @@ impl Store {
let mut next_hashes = Vec::with_capacity(models.len());
let mut unique_hashes = HashSet::with_capacity(models.len());
for model in &models {
let hash = model_hash(&base_url, model.endpoint_type, &model.model_id)?;
let hash = model_hash(&base_url, api_key, model.endpoint_type, &model.model_id)?;
if !unique_hashes.insert(hash.clone()) {
return Err(Error::Config(format!(
"8-character model hash collision: {hash}"
@@ -271,6 +276,7 @@ impl Store {
for input in inputs {
let hash = model_hash(
&provider.endpoint.base_url,
&provider.api_key,
input.endpoint_type,
&input.model_id,
)?;
@@ -327,6 +333,7 @@ impl Store {
.expect("model provider must exist");
let next_hash = model_hash(
&provider.endpoint.base_url,
&provider.api_key,
input.endpoint_type,
&input.model_id,
)?;
@@ -680,6 +687,33 @@ mod tests {
assert_eq!(store.provider_models(false).await.unwrap().len(), 2);
}
#[tokio::test]
async fn allows_same_endpoint_and_model_with_different_api_keys() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("credential-models.db").display()
))
.await
.unwrap();
let first_provider = provider();
let mut second_provider = provider();
second_provider.name = "Second".into();
second_provider.api_key = Some("different-secret".into());
let (_, first_model) = store
.create_provider_with_model(&first_provider, &model("model-a"))
.await
.unwrap();
let (_, second_model) = store
.create_provider_with_model(&second_provider, &model("model-a"))
.await
.unwrap();
assert_ne!(first_model.model_hash, second_model.model_hash);
assert_eq!(store.provider_models(false).await.unwrap().len(), 2);
}
#[tokio::test]
async fn adds_multiple_models_to_existing_provider_atomically() {
let directory = tempfile::tempdir().unwrap();
@@ -750,6 +784,55 @@ mod tests {
models[0].model_hash,
model_hash(
&updated_provider.base_url,
input.api_key.as_deref().unwrap(),
models[0].endpoint_type,
&models[0].model_id,
)
.unwrap()
);
let detached: Option<String> =
sqlx::query_scalar("SELECT model_hash FROM llm_calls WHERE call_id = ?")
.bind("call-1")
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(detached, None);
}
#[tokio::test]
async fn updating_provider_api_key_rehashes_its_models() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("provider-key-update.db").display()
))
.await
.unwrap();
let (created_provider, original) = store
.create_provider_with_model(&provider(), &model("model-a"))
.await
.unwrap();
insert_call(&store, &created_provider, &original).await;
let mut input = provider();
input.api_key = Some("different-secret".into());
store
.update_provider(created_provider.provider_id, &input)
.await
.unwrap();
assert!(store
.provider_model(&original.model_hash)
.await
.unwrap()
.is_none());
let models = store.provider_models(false).await.unwrap();
assert_eq!(models.len(), 1);
assert_eq!(
models[0].model_hash,
model_hash(
&created_provider.base_url,
"different-secret",
models[0].endpoint_type,
&models[0].model_id,
)
+1 -1
View File
@@ -77,7 +77,7 @@ async fn provider_secret_is_write_only_and_model_hash_is_stable() {
)
.await
.unwrap();
assert_eq!(model.model_hash, "f246010a");
assert_eq!(model.model_hash, "bab5019a");
assert!(model.supports_image_generation);
}
+46 -11
View File
@@ -148,6 +148,35 @@ async fn duplicate_usage_is_rejected_instead_of_guessing_which_total_is_final()
assert!(matches!(failure.failure, RunFailure::Protocol(_)));
}
#[tokio::test]
async fn duplicate_tool_call_ids_are_rejected_across_distinct_indexes() {
let (sender, _receiver) = tokio::sync::mpsc::channel(8);
let failure = consume_model_cycle(
provider_stream(vec![
ModelEvent::Start {
model_call_id: "model-call".into(),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::ToolCallStart {
index: 1,
call_id: "call-1".into(),
name: "Read".into(),
},
]),
&sender,
&CancellationToken::new(),
)
.await
.unwrap_err();
assert!(matches!(failure.failure, RunFailure::Protocol(_)));
}
#[tokio::test]
async fn openai_chat_raw_stream_and_request_projection_match_the_endpoint() {
let (base_url, mut requests, server) = fixture_server(
@@ -457,12 +486,15 @@ async fn openai_responses_preserves_delta_that_repeats_the_streamed_suffix() {
}
#[tokio::test]
async fn openai_responses_completed_object_recovers_missing_item_events() {
async fn openai_responses_completed_snapshot_does_not_reindex_streamed_tool() {
let (base_url, _requests, server) = fixture_server(
"/v1/responses",
concat!(
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"reasoning\",\"id\":\"reasoning-1\",\"encrypted_content\":\"opaque\"}}\n\n",
"data: {\"type\":\"response.output_item.added\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\"}}\n\n",
"data: {\"type\":\"response.function_call_arguments.done\",\"output_index\":1,\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}\n\n",
"data: {\"type\":\"response.output_item.done\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}}\n\n",
"data: {\"type\":\"response.completed\",\"response\":{\"output\":[",
"{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"ok\"}]},",
"{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}",
"]}}\n\n",
),
@@ -472,18 +504,21 @@ async fn openai_responses_completed_object_recovers_missing_item_events() {
reqwest::Client::new(),
config(ProviderKind::OpenAiResponses, base_url, None),
);
let (sender, _receiver) = tokio::sync::mpsc::channel(32);
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
let cycle = consume_model_cycle(
provider.stream(invocation(), CancellationToken::new()),
&sender,
&CancellationToken::new(),
)
.await
.unwrap();
server.abort();
assert!(events
.iter()
.any(|event| matches!(event, ModelEvent::TextDelta(text) if text == "ok")));
assert!(events.iter().any(|event| matches!(event, ModelEvent::ToolCallStart { call_id, name, .. } if call_id == "call-1" && name == "Read")));
assert_eq!(
events.last(),
Some(&ModelEvent::Done(FinishReason::ToolUse))
);
assert_eq!(cycle.calls.len(), 1);
assert_eq!(cycle.calls[0].index, 1);
assert_eq!(cycle.calls[0].call_id, "call-1");
assert_eq!(cycle.calls[0].arguments["path"], "a");
}
#[tokio::test]