mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 03:56:45 +08:00
fix: harden provider, tool, and desktop behavior
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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<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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user