diff --git a/Makefile b/Makefile index e8e11df..b23ddc2 100644 --- a/Makefile +++ b/Makefile @@ -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 . diff --git a/apps/desktop/src/components/virtual/VirtualListEngine.tsx b/apps/desktop/src/components/virtual/VirtualListEngine.tsx index 533d1b0..839d45d 100644 --- a/apps/desktop/src/components/virtual/VirtualListEngine.tsx +++ b/apps/desktop/src/components/virtual/VirtualListEngine.tsx @@ -174,6 +174,7 @@ export function VirtualList(props: VirtualListProps) { const shouldResetScrollRef = useRef(false) const scrollApiRef = useRef(null) const scrollStateRef = useRef(null) + const contentElementRef = useRef(null) const spacerRef = useRef(null) const [contentInsets, setContentInsets] = useState({ top: 0, @@ -192,7 +193,8 @@ export function VirtualList(props: VirtualListProps) { }) 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(props: VirtualListProps) { ) }, []) + 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) { diff --git a/server/src/cursor/tools/result/exec/mod.rs b/server/src/cursor/tools/result/exec/mod.rs index 06cfec9..c4d52aa 100644 --- a/server/src/cursor/tools/result/exec/mod.rs +++ b/server/src/cursor/tools/result/exec/mod.rs @@ -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 { 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); } diff --git a/server/src/cursor/tools/result/gate.rs b/server/src/cursor/tools/result/gate.rs new file mode 100644 index 0000000..57cc5a6 --- /dev/null +++ b/server/src/cursor/tools/result/gate.rs @@ -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); + } +} diff --git a/server/src/cursor/tools/result/mod.rs b/server/src/cursor/tools/result/mod.rs index 369a03d..4a8add4 100644 --- a/server/src/cursor/tools/result/mod.rs +++ b/server/src/cursor/tools/result/mod.rs @@ -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 { diff --git a/server/src/model/provider.rs b/server/src/model/provider.rs index 8669271..2295cfe 100644 --- a/server/src/model/provider.rs +++ b/server/src/model/provider.rs @@ -138,7 +138,12 @@ pub fn normalize_base_url(value: &str) -> Result { Ok(url.as_str().trim_end_matches('/').to_string()) } -pub fn model_hash(base_url: &str, provider_type: ProviderType, model_id: &str) -> Result { +pub fn model_hash( + base_url: &str, + api_key: &str, + provider_type: ProviderType, + model_id: &str, +) -> Result { 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() ); } diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 788256a..d1b4578 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -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) { diff --git a/server/src/store/providers.rs b/server/src/store/providers.rs index 2deef4c..e899321 100644 --- a/server/src/store/providers.rs +++ b/server/src/store/providers.rs @@ -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(¤t.api_key); let custom_headers = merge_custom_headers(¤t.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 = + 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, ) diff --git a/server/tests/provider_console.rs b/server/tests/provider_console.rs index 7cd1567..461ace9 100644 --- a/server/tests/provider_console.rs +++ b/server/tests/provider_console.rs @@ -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); } diff --git a/server/tests/provider_stream.rs b/server/tests/provider_stream.rs index 0fc7189..24ebf3d 100644 --- a/server/tests/provider_stream.rs +++ b/server/tests/provider_stream.rs @@ -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]