Merge branch 'main' of github.com:leookun/cursor-byok into fix/pr343-minimal-compaction-prompt

# Conflicts:
#	server/src/config.rs
This commit is contained in:
leookun
2026-08-27 21:49:59 +08:00
96 changed files with 14712 additions and 733 deletions
+1 -1
View File
@@ -2,7 +2,7 @@
"tools": [
"Shell", "Grep", "Delete", "WebSearch", "WebFetch", "GenerateImage",
"EditNotebook", "TodoWrite", "StrReplace", "Write", "Read", "ReadLints",
"Glob", "AskQuestion", "Task", "AwaitShell", "GetMcpTools",
"Glob", "AskQuestion", "Task", "GetMcpTools",
"FetchMcpResource", "SwitchMode", "CallMcpTool", "SembleSearch",
"SembleFindRelated"
]
+1 -1
View File
@@ -2,7 +2,7 @@
"tools": [
"Shell", "Grep", "Delete", "WebSearch", "WebFetch", "GenerateImage",
"ReadLints", "EditNotebook", "TodoWrite", "StrReplace", "Write", "Read",
"Glob", "AwaitShell", "GetMcpTools", "FetchMcpResource", "SwitchMode",
"Glob", "GetMcpTools", "FetchMcpResource", "SwitchMode",
"UpdateCurrentStep", "CallMcpTool", "SembleSearch", "SembleFindRelated"
]
}
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -11,7 +11,7 @@ pub struct MessageInsertion {
#[derive(Debug)]
pub enum ClientCommand {
ToolResult(ToolResult),
RuntimeMessage(CanonicalMessage),
InterruptWithMessage(CanonicalMessage),
RuntimeEvent(RuntimeEvent),
InsertMessages(MessageInsertion),
ClientClosed { error: String },
+1 -1
View File
@@ -9,7 +9,7 @@ use crate::run::RunOutcome;
pub enum CommitCause {
InitialMessages,
ToolRoundStarted(ToolRoundId),
ToolResult { call_id: String },
ToolResult { call_id: String, interrupted: bool },
FinalTurn,
Compaction { summary: String },
RuntimeEvent { event_id: String },
+3 -2
View File
@@ -10,6 +10,7 @@ const DATABASE_FILE_NAME: &str = "cursor-byok.db";
const V0049_DATA_DIR_NAME: &str = ".cursor-local-assistant-v2";
const V0049_CONFIG_FILE_NAME: &str = "config.yaml";
const COMPACTION_PROMPT_PATH: &str = "prompts/compaction.md";
const DEFAULT_PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(3000);
pub fn managed_data_dir() -> Result<PathBuf> {
let home_dir = dirs::home_dir()
@@ -91,7 +92,7 @@ impl Config {
Ok(value) => Duration::from_secs(value.parse().map_err(|error| {
Error::Config(format!("invalid CURSOR_PROVIDER_TIMEOUT_SECONDS: {error}"))
})?),
Err(env::VarError::NotPresent) => Duration::from_secs(300),
Err(env::VarError::NotPresent) => DEFAULT_PROVIDER_REQUEST_TIMEOUT,
Err(error) => {
return Err(Error::Config(format!(
"invalid CURSOR_PROVIDER_TIMEOUT_SECONDS: {error}"
@@ -132,7 +133,7 @@ impl Config {
.parse()
.expect("desktop listen address is static"),
database_url: default_database_url()?,
provider_request_timeout: Duration::from_secs(300),
provider_request_timeout: DEFAULT_PROVIDER_REQUEST_TIMEOUT,
console: None,
use_persisted_ports: true,
})
+4
View File
@@ -285,6 +285,10 @@ impl CursorActor {
{
continue;
}
if tool_runtime.is_interrupted(throw.id).await {
tool_runtime.discard_exec(throw.id).await;
continue;
}
match tool_runtime.take_exec(throw.id).await {
Some(pending) => results_tx.send_error(
crate::Error::Protocol(format!(
+13
View File
@@ -152,6 +152,19 @@ pub fn context_injection_queued(injection_id: String) -> pb::AgentServerMessage
))
}
pub fn context_injection_rejected(injection_id: String, reason: String) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ContextInjectionState(
pb::ContextInjectionStateUpdate {
injection_id,
state: Some(pb::ContextInjectionState {
state: Some(pb::context_injection_state::State::Rejected(
pb::ContextInjectionRejected { reason },
)),
}),
},
))
}
pub fn context_injection_delivered(
injection_id: String,
delivery_batch_id: String,
-12
View File
@@ -188,7 +188,6 @@ pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
"updatecurrentstep" => {
Tool::CommunicateUpdateToolCall(pb::CommunicateUpdateToolCall::default())
}
"awaitshell" => Tool::AwaitToolCall(pb::AwaitToolCall::default()),
"getmcptools" => Tool::GetMcpToolsToolCall(pb::GetMcpToolsToolCall::default()),
_ => return Err(Error::Protocol(format!("unsupported tool: {name}"))),
};
@@ -457,17 +456,6 @@ pub fn render_tool_call(call: &ToolCall, completed: bool) -> Result<pb::ToolCall
chars: string("chars"),
})
}
Some(pb::tool_call::Tool::AwaitToolCall(tool)) => {
tool.args = Some(pb::AwaitArgs {
task_id: string("shell_id"),
block_until_ms: call
.arguments
.get("block_until_ms")
.and_then(Value::as_u64)
.map(|v| v as u32),
regex: optional("pattern"),
})
}
Some(pb::tool_call::Tool::GetMcpToolsToolCall(tool)) => {
tool.args = Some(pb::GetMcpToolsArgs {
server: optional("server"),
-1
View File
@@ -158,7 +158,6 @@ fn tool_identifier(name: &str, dynamic_tools: &HashSet<String>) -> String {
return name.into();
}
match name {
"AwaitShell" => "AWAIT".into(),
"CallMcpTool" | "SembleSearch" | "SembleFindRelated" => "MCP".into(),
"CreatePlan" => "CREATE_PLAN_V2".into(),
"UpdateCurrentStep" => "COMMUNICATE_UPDATE".into(),
+5
View File
@@ -505,6 +505,11 @@ fn normalize_mcp_parameters(tool_name: &str, mut parameters: Value) -> Result<Va
if !object_only_union {
return Err(invalid_mcp_parameters(tool_name));
}
// OpenAI-compatible function schemas (and the corresponding schema
// validators used by other providers) require the root schema to declare
// an object type. Cursor's app-control MCP sometimes sends an object-only
// `anyOf`/`oneOf` schema without that root annotation. Preserve the union
// while adding the annotation to the model-facing copy.
schema.insert("type".into(), Value::String("object".into()));
Ok(parameters)
}
+92 -28
View File
@@ -122,6 +122,9 @@ impl CursorSession {
let mut response_text = String::new();
let mut response_thinking = String::new();
let mut active_round = None::<ToolRoundId>;
let mut active_tool_calls = HashSet::<String>::new();
let mut interrupted_rounds = HashSet::<ToolRoundId>::new();
let mut interrupted_tool_calls = HashSet::<String>::new();
let mut final_checkpoint = None::<FinalCheckpoints>;
let mut compaction_checkpoint = None::<pb::ConversationStateStructure>;
let mut turn_usage = None::<Usage>;
@@ -130,13 +133,16 @@ impl CursorSession {
let mut presentation = Presentation::default();
loop {
let input = if let Some(completion) = ready.pop_front() {
let input = if let Ok(action) = self.runtime_actions.try_recv() {
Input::RuntimeAction(Some(Box::new(action)))
} else if let Some(completion) = ready.pop_front() {
Input::Completion(completion)
} else {
tokio::select! {
biased;
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
event = self.core.events.recv() => Input::Event(event),
completion = self.results.recv() => Input::CompletionResult(completion),
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
failure = worker.failures.recv(), if checkpoint_worker_open => Input::CheckpointFailure(failure),
}
};
@@ -147,15 +153,16 @@ impl CursorSession {
}
Input::Completion(completion) => {
if let Some(completion) = self
.forward_completion(completion, &mut completions)
.forward_completion(completion, &mut completions, &interrupted_tool_calls)
.await?
{
ready.push_back(completion);
}
}
Input::CompletionResult(Some(result)) => {
if let Some(completion) =
self.forward_completion(result?, &mut completions).await?
if let Some(completion) = self
.forward_completion(result?, &mut completions, &interrupted_tool_calls)
.await?
{
ready.push_back(completion);
}
@@ -164,7 +171,15 @@ impl CursorSession {
return Err(Error::Protocol("tool result channel closed".into()));
}
Input::RuntimeAction(Some(action)) => {
self.forward_injection(*action).await?;
self.forward_injection(
*action,
active_round.as_ref(),
&active_tool_calls,
&completions,
&mut interrupted_rounds,
&mut interrupted_tool_calls,
)
.await?;
}
Input::RuntimeAction(None) => {
return Err(Error::Protocol("runtime action channel closed".into()));
@@ -283,7 +298,23 @@ impl CursorSession {
round_id,
calls: round_calls,
} => {
active_round = Some(round_id);
active_round = Some(round_id.clone());
active_tool_calls = round_calls
.iter()
.map(|call| call.call_id.clone())
.collect();
// Runtime actions are deliberately prioritized over core events. An
// injection can therefore be observed before the already-queued
// ToolRoundStarted event reaches this session. In that case the
// accepted injection is still pending delivery and this round must be
// detached without starting any root tools.
if interrupted_rounds.contains(&round_id)
|| !self.pending_injections.is_empty()
{
interrupted_rounds.insert(round_id.clone());
interrupted_tool_calls.extend(active_tool_calls.iter().cloned());
continue;
}
for dispatched in self
.tools
.start_batch(
@@ -348,12 +379,11 @@ impl CursorSession {
active_round = Some(round_id.clone());
}
let mut tool_round_settled = false;
if let CommitCause::ToolResult { call_id } = &state.cause {
let completion = completions.remove(call_id).ok_or_else(|| {
Error::Protocol(format!(
"core committed a tool result without typed Cursor state: {call_id}"
))
})?;
if let CommitCause::ToolResult {
call_id,
interrupted,
} = &state.cause
{
let snapshot = self
.store
.tool_round(active_round.as_ref().ok_or_else(|| {
@@ -372,9 +402,16 @@ impl CursorSession {
"committed call is absent from tool round: {call_id}"
))
})?;
self.handle
.emit(&interaction::tool_completed(call, &completion))?;
presentation.tool_completed(&completion);
if !interrupted {
let completion = completions.remove(call_id).ok_or_else(|| {
Error::Protocol(format!(
"core committed a tool result without typed Cursor state: {call_id}"
))
})?;
self.handle
.emit(&interaction::tool_completed(call, &completion))?;
presentation.tool_completed(&completion);
}
completed.insert(call_id.clone());
tool_round_settled = snapshot.status == ToolRoundStatus::Settled;
}
@@ -490,7 +527,10 @@ impl CursorSession {
return Err(error);
}
}
active_round = None;
if let Some(round_id) = active_round.take() {
interrupted_rounds.remove(&round_id);
}
active_tool_calls.clear();
self.tool_runtime.clear_completed().await;
} else if !matches!(&state.cause, CommitCause::ToolResult { .. })
&& active_round.is_some()
@@ -603,7 +643,11 @@ impl CursorSession {
&self,
mut completion: ToolCompletion,
completions: &mut HashMap<String, ToolCompletion>,
interrupted_tool_calls: &HashSet<String>,
) -> Result<Option<ToolCompletion>> {
if interrupted_tool_calls.contains(&completion.result().call_id) {
return Ok(None);
}
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)?;
@@ -635,19 +679,33 @@ impl CursorSession {
Ok(dispatched.completion)
}
async fn forward_injection(&mut self, action: pb::InjectContextAction) -> Result<()> {
async fn forward_injection(
&mut self,
action: pb::InjectContextAction,
active_round: Option<&ToolRoundId>,
active_tool_calls: &HashSet<String>,
completions: &HashMap<String, ToolCompletion>,
interrupted_rounds: &mut HashSet<ToolRoundId>,
interrupted_tool_calls: &mut HashSet<String>,
) -> Result<()> {
if action.injection_id.is_empty() {
return Err(Error::Protocol(
"InjectContextAction has no injection_id".into(),
));
}
if self.injection_ids.contains(&action.injection_id) {
return Ok(());
}
if action.expected_run_id != self.context.request_id {
return Err(Error::Protocol(format!(
let reason = format!(
"InjectContextAction expected run {}, active run is {}",
action.expected_run_id, self.context.request_id
)));
}
if self.injection_ids.contains(&action.injection_id) {
);
self.handle.emit(&interaction::context_injection_rejected(
action.injection_id.clone(),
reason,
))?;
self.injection_ids.insert(action.injection_id);
return Ok(());
}
let user_message = match action.payload.as_ref() {
@@ -675,25 +733,31 @@ impl CursorSession {
);
self.handle
.emit(&interaction::context_injection_queued(injection_id.clone()))?;
interrupted_tool_calls.extend(
active_tool_calls
.iter()
.filter(|call_id| !completions.contains_key(*call_id))
.cloned(),
);
if let Some(round_id) = active_round {
interrupted_rounds.insert(round_id.clone());
}
self.interrupt_execs().await;
if self
.core
.commands
.send(ClientCommand::RuntimeMessage(message))
.send(ClientCommand::InterruptWithMessage(message))
.await
.is_err()
{
self.pending_injections.remove(&injection_id);
return Err(Error::RunNotFound(self.context.request_id.clone()));
}
self.interrupt_execs().await;
Ok(())
}
async fn interrupt_execs(&self) {
// Keep runtime entries until Cursor returns the aborted result. The core tool
// round needs that terminal result before it can append the injected context
// after the complete assistant/tool pair and continue the same Run.
for id in self.tool_runtime.running_exec_ids().await {
for id in self.tools.interrupt_for_message().await {
let _ = self.handle.emit(&codec::abort(id));
}
}
+1 -3
View File
@@ -2,7 +2,5 @@ mod request;
mod response;
pub use request::{abort, mcp_request, mcp_state_request, request};
pub(crate) use request::{
await_read_request, edit_read_request, json_object_to_prost, mcp_meta_request,
};
pub(crate) use request::{edit_read_request, json_object_to_prost, mcp_meta_request};
pub use response::{client_event, stream_closed, ClientExecEvent};
-26
View File
@@ -191,32 +191,6 @@ pub(crate) fn edit_read_request(id: u32, call: &ToolCall) -> Result<pb::AgentSer
))
}
pub(crate) fn await_read_request(
id: u32,
call: &ToolCall,
context: &ExecContext,
) -> Result<pb::AgentServerMessage> {
let task_id = call
.arguments
.get("shell_id")
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol("AwaitShell is missing shell_id".into()))?;
Ok(server_message(
id,
call,
pb::exec_server_message::Message::ReadArgs(pb::ReadArgs {
path: format!(
"{}/{}.txt",
context.terminals_folder.trim_end_matches('/'),
task_id
),
tool_call_id: call.call_id.clone(),
..Default::default()
}),
Some(false),
))
}
pub(super) fn edit_write_request(
id: u32,
call: &ToolCall,
+24 -74
View File
@@ -12,7 +12,7 @@ use crate::{
Error, Result,
};
use super::request::{await_read_request, edit_write_request};
use super::request::edit_write_request;
pub enum ClientExecEvent {
Delta(Box<pb::AgentServerMessage>),
@@ -25,6 +25,12 @@ pub async fn client_event(
message: &pb::ExecClientMessage,
pending: &CursorToolRuntime,
) -> Result<ClientExecEvent> {
if pending.is_interrupted(message.id).await {
if message.message.as_ref().is_some_and(is_terminal) {
pending.discard_exec(message.id).await;
}
return Ok(ClientExecEvent::Pending);
}
let call = match pending.exec_call(message.id).await {
Some(call) => call,
None if pending.completed_call(message.id).await.is_some() => {
@@ -47,7 +53,6 @@ pub async fn client_event(
let entry = take(message.id, pending).await?;
return match entry.stage {
ExecStage::EditRead => advance_edit(entry, wire_result, pending).await,
ExecStage::Await(_) => advance_await(entry, wire_result, pending).await,
ExecStage::Direct | ExecStage::DynamicMcp(_) | ExecStage::EditWrite(_) => {
completed(entry, wire_result.clone())
}
@@ -131,6 +136,10 @@ pub async fn client_event(
}
pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Option<ToolCompletion>> {
if pending.is_interrupted(id).await {
pending.discard_exec(id).await;
return Ok(None);
}
let Some(entry) = pending.take_exec(id).await else {
return Ok(None);
};
@@ -177,79 +186,20 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
)?))
}
async fn advance_await(
entry: PendingExec,
result: &pb::exec_client_message::Message,
registry: &CursorToolRuntime,
) -> Result<ClientExecEvent> {
let read = match result {
pb::exec_client_message::Message::ReadResult(result)
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
_ => return Err(Error::Protocol("AwaitShell expected ReadResult".into())),
};
let ExecStage::Await(state) = &entry.stage else {
return Err(Error::Protocol(
"AwaitShell result reached a non-await execution stage".into(),
));
};
let content = match read.result.as_ref() {
Some(pb::read_result::Result::Success(success)) => match success.output.as_ref() {
Some(pb::read_success::Output::Content(content)) => content.as_str(),
_ => "",
},
Some(pb::read_result::Result::FileNotFound(_)) => "",
Some(pb::read_result::Result::Error(error)) => {
return Ok(ClientExecEvent::Completed(Box::new(result::await_error(
entry,
&error.error,
)?)))
}
_ => "",
};
let regex_match = state
.regex
.as_ref()
.map(|pattern| regex::Regex::new(pattern))
.transpose()
.map_err(|error| Error::Protocol(format!("invalid AwaitShell pattern: {error}")))?
.and_then(|pattern| {
pattern
.find(content)
.map(|found| found.as_str().to_string())
});
let exit_code = content.lines().find_map(|line| {
line.strip_prefix("exit_code:")
.and_then(|value| value.trim().parse::<i32>().ok())
});
if regex_match.is_some() || exit_code.is_some() || std::time::Instant::now() >= state.deadline {
return Ok(ClientExecEvent::Completed(Box::new(result::await_result(
entry,
content.len() as u64,
regex_match,
exit_code,
)?)));
fn is_terminal(message: &pb::exec_client_message::Message) -> bool {
use pb::{exec_client_message::Message, shell_stream::Event};
match message {
Message::ShellStream(stream) => matches!(
stream.event.as_ref(),
Some(Event::Exit(_))
| Some(Event::Backgrounded(_))
| Some(Event::Rejected(_))
| Some(Event::PermissionDenied(_))
| Some(Event::SandboxUnsupported(_))
),
_ => true,
}
let state = match entry.stage {
ExecStage::Await(state) => state,
_ => {
return Err(Error::Protocol(
"AwaitShell result changed execution stage".into(),
))
}
};
let wait = state
.deadline
.saturating_duration_since(std::time::Instant::now())
.min(std::time::Duration::from_secs(1));
tokio::time::sleep(wait).await;
let call = entry.call.clone();
let context = entry.context.clone();
let id = registry
.reserve_await_again(&call, &context, state, entry.started_at_ms)
.await?;
Ok(ClientExecEvent::Message(Box::new(await_read_request(
id, &call, &context,
)?)))
}
async fn advance_edit(
@@ -1,54 +0,0 @@
//! AwaitShell's timed and file-backed execution paths.
use crate::{model::ToolCall, Error, Result};
use super::ToolStart;
use crate::cursor::tools::{
codec, result,
result::ToolResultSender,
runtime::{CursorToolRuntime, ExecContext},
};
pub(super) async fn start(
runtime: &CursorToolRuntime,
results: &ToolResultSender,
call: &ToolCall,
context: &ExecContext,
) -> Result<ToolStart> {
let message = if call
.arguments
.get("shell_id")
.and_then(serde_json::Value::as_str)
.is_some()
{
let id = runtime.reserve_await(call, context).await?;
Some(codec::await_read_request(id, call, context)?)
} else {
wait_without_shell_id(results, call)?;
None
};
Ok(ToolStart {
messages: message.into_iter().collect(),
completion: None,
})
}
fn wait_without_shell_id(results: &ToolResultSender, call: &ToolCall) -> Result<()> {
let block_ms = call
.arguments
.get("block_until_ms")
.and_then(serde_json::Value::as_u64)
.unwrap_or(30_000);
if block_ms == 0 || block_ms > 7_140_000 {
return Err(Error::Protocol(
"AwaitShell without shell_id requires block_until_ms in 1..=7140000".into(),
));
}
let call = call.clone();
let results = results.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(block_ms)).await;
results.send(result::await_sleep(&call, block_ms));
});
Ok(())
}
+110 -2
View File
@@ -1,4 +1,3 @@
mod await_shell;
mod edit;
mod exec;
mod interaction;
@@ -51,6 +50,9 @@ pub(super) async fn start(
return local::subagents_disabled(call);
}
let normalized_call = normalize_block_until_ms(call)?;
let call = normalized_call.as_ref().unwrap_or(call);
match normalized(&call.name).as_str() {
"shell" | "read" | "delete" | "grep" | "glob" | "readlints" | "task" | "callmcptool"
| "fetchmcpresource" | "getmcptools" => exec::start(runtime, call, context).await,
@@ -58,12 +60,60 @@ pub(super) async fn start(
"askquestion" | "websearch" | "webfetch" | "switchmode" | "createplan"
| "generateimage" => interaction::start(runtime, call).await,
"todowrite" | "updatecurrentstep" => local::start(call, message_index),
"awaitshell" => await_shell::start(runtime, results, call, context).await,
"semblesearch" | "semblefindrelated" => semble::start(results, call, store.cloned()),
_ => Err(Error::Protocol(format!("unsupported tool: {}", call.name))),
}
}
fn normalize_block_until_ms(call: &ToolCall) -> Result<Option<ToolCall>> {
if normalized(&call.name) != "shell" {
return Ok(None);
}
let Some(value) = call.arguments.get("block_until_ms") else {
return Ok(None);
};
let integer = if let Some(value) = value.as_i64() {
value
} else {
let value = value.as_f64().ok_or_else(|| {
Error::Protocol(format!("{} block_until_ms must be an integer", call.name))
})?;
if !value.is_finite() || value.fract() != 0.0 {
return Err(Error::Protocol(format!(
"{} block_until_ms must be an integer",
call.name
)));
}
if value < i64::MIN as f64 || value > i64::MAX as f64 {
return Err(Error::Protocol(format!(
"{} block_until_ms is out of range",
call.name
)));
}
value as i64
};
if integer < 0 {
return Err(Error::Protocol(format!(
"{} block_until_ms is out of range",
call.name
)));
}
if value.as_i64().is_some() {
return Ok(None);
}
let mut normalized_call = call.clone();
normalized_call
.arguments
.as_object_mut()
.ok_or_else(|| Error::Protocol(format!("{} arguments must be a JSON object", call.name)))?
.insert("block_until_ms".into(), serde_json::Value::from(integer));
Ok(Some(normalized_call))
}
fn is_mcp_auth(call: &ToolCall) -> bool {
normalized(&call.name) == "callmcptool"
&& call
@@ -89,3 +139,61 @@ pub(super) fn normalized(name: &str) -> String {
.flat_map(char::to_lowercase)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn tool(name: &str, arguments: serde_json::Value) -> ToolCall {
ToolCall {
index: 0,
call_id: "call-1".into(),
model_call_id: "model-call-1".into(),
name: name.into(),
arguments_text: arguments.to_string(),
arguments,
}
}
#[test]
fn shell_accepts_integer_valued_float_timeout() {
let call = tool(
"Shell",
serde_json::json!({"command": "echo ok", "block_until_ms": 45_000.0}),
);
let call = normalize_block_until_ms(&call).unwrap().unwrap();
assert_eq!(call.arguments["block_until_ms"].as_i64(), Some(45_000));
}
#[test]
fn shell_rejects_fractional_timeout() {
let call = tool(
"Shell",
serde_json::json!({"command": "echo ok", "block_until_ms": 30_000.5}),
);
let error = normalize_block_until_ms(&call).unwrap_err();
assert_eq!(
error.to_string(),
"protocol error: Shell block_until_ms must be an integer"
);
}
#[test]
fn shell_rejects_negative_timeout_instead_of_defaulting() {
let call = tool(
"Shell",
serde_json::json!({"command": "echo ok", "block_until_ms": -1}),
);
let error = normalize_block_until_ms(&call).unwrap_err();
assert_eq!(
error.to_string(),
"protocol error: Shell block_until_ms is out of range"
);
}
}
+8
View File
@@ -155,6 +155,11 @@ impl ToolDispatcher {
.map(Some)
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
self.edit_schedule.lock().await.clear();
self.runtime.interrupt_for_message().await
}
async fn start(
&self,
call: &ToolCall,
@@ -193,6 +198,9 @@ impl ToolDispatcher {
&self,
response: &pb::InteractionResponse,
) -> Result<ClientToolEvent> {
if self.runtime.is_interrupted(response.id).await {
return Ok(ClientToolEvent::Pending);
}
let pending = match self.runtime.take_interaction(response.id).await {
Some(pending) => pending,
None if self.runtime.completed_call(response.id).await.is_some() => {
@@ -1,145 +0,0 @@
use serde_json::Value;
use crate::{
cursor::proto::agent::v1 as pb,
model::{ToolCall, ToolResult},
Error, Result,
};
use super::{now_ms, ToolCompletion};
use crate::cursor::tools::runtime::{ExecStage, PendingExec};
pub(crate) fn await_result(
pending: PendingExec,
output_length: u64,
regex_match: Option<String>,
exit_code: Option<i32>,
) -> Result<ToolCompletion> {
let ExecStage::Await(state) = &pending.stage else {
return Err(Error::Protocol(
"AwaitShell completion reached a non-await execution stage".into(),
));
};
let runtime_ms = now_ms().saturating_sub(pending.started_at_ms);
let result = if exit_code.is_some() {
pb::await_success::AwaitResult::Complete(pb::AwaitTaskComplete {
task_id: state.task_id.clone(),
runtime_ms,
output_file_path: state.output_file_path.clone(),
output_length,
regex_requested: state.regex.is_some(),
regex_match,
exit_code,
wake_reason: Some("task_complete".into()),
})
} else {
pb::await_success::AwaitResult::StillRunning(pb::AwaitTaskStillRunning {
task_id: state.task_id.clone(),
runtime_ms,
output_file_path: state.output_file_path.clone(),
output_length,
regex_requested: state.regex.is_some(),
regex_match,
wake_reason: Some("timeout_or_pattern".into()),
})
};
let content = serde_json::json!({
"task_id": state.task_id,
"output_file_path": state.output_file_path,
"output_length": output_length,
"exit_code": exit_code,
})
.to_string();
completion(
&pending,
content,
false,
pb::await_result::Result::Success(pb::AwaitSuccess {
await_result: Some(result),
}),
)
}
pub(crate) fn await_error(pending: PendingExec, error: &str) -> Result<ToolCompletion> {
completion(
&pending,
error.into(),
true,
pb::await_result::Result::Error(pb::AwaitError {
error: error.into(),
}),
)
}
fn completion(
pending: &PendingExec,
content: String,
is_error: bool,
result: pb::await_result::Result,
) -> Result<ToolCompletion> {
let ExecStage::Await(state) = &pending.stage else {
return Err(Error::Protocol(
"AwaitShell completion reached a non-await execution stage".into(),
));
};
Ok(ToolCompletion::new(
&pending.call,
pending.started_at_ms,
ToolResult {
call_id: pending.call.call_id.clone(),
content,
is_error,
image: None,
},
pb::tool_call::Tool::AwaitToolCall(pb::AwaitToolCall {
args: Some(pb::AwaitArgs {
task_id: state.task_id.clone(),
block_until_ms: pending
.call
.arguments
.get("block_until_ms")
.and_then(Value::as_u64)
.map(|value| value as u32),
regex: state.regex.clone(),
}),
result: Some(pb::AwaitResult {
result: Some(result),
}),
}),
))
}
pub(crate) fn await_sleep(call: &ToolCall, runtime_ms: u64) -> ToolCompletion {
ToolCompletion::new(
call,
now_ms().saturating_sub(runtime_ms),
ToolResult {
call_id: call.call_id.clone(),
content: format!("Waited {runtime_ms} ms"),
is_error: false,
image: None,
},
pb::tool_call::Tool::AwaitToolCall(pb::AwaitToolCall {
args: Some(pb::AwaitArgs {
task_id: String::new(),
block_until_ms: Some(runtime_ms as u32),
regex: None,
}),
result: Some(pb::AwaitResult {
result: Some(pb::await_result::Result::Success(pb::AwaitSuccess {
await_result: Some(pb::await_success::AwaitResult::StillRunning(
pb::AwaitTaskStillRunning {
task_id: String::new(),
runtime_ms,
output_file_path: String::new(),
output_length: 0,
regex_requested: false,
regex_match: None,
wake_reason: Some("sleep_complete".into()),
},
)),
})),
}),
}),
)
}
+818 -49
View File
@@ -1,13 +1,54 @@
use crate::cursor::proto::agent::v1 as pb;
use std::collections::BTreeMap;
use crate::{cursor::proto::agent::v1 as pb, model::limit_tool_result_text};
const KIB: usize = 1024;
const READ_CONTENT_LIMIT: usize = 64 * KIB;
const READ_BINARY_LIMIT: usize = 32 * KIB;
const SHELL_STREAM_LIMIT: usize = 16 * KIB;
const SHELL_CONTENT_LIMIT: usize = 32 * KIB;
const SHELL_INTERLEAVED_LIMIT: usize = 32 * KIB;
const GREP_CONTENT_LIMIT: usize = 32 * KIB;
const GREP_MATCH_LIMIT: usize = 2 * KIB;
const GREP_MATCHES_PER_FILE: usize = 100;
const GREP_TOTAL_MATCHES: usize = 300;
const GREP_LIST_LIMIT: usize = 300;
const GLOB_FILE_LIMIT: usize = 200;
const EDIT_RESULT_LIMIT: usize = 32 * KIB;
const PATCH_EDIT_RESULT_LIMIT: usize = 4 * KIB;
const MCP_TEXT_LIMIT: usize = 32 * KIB;
const MCP_CONTENT_ITEM_LIMIT: usize = 20;
const MCP_STRUCTURED_LIMIT: usize = 32 * KIB;
const MCP_BINARY_LIMIT: usize = 32 * KIB;
const MCP_RESOURCE_LIMIT: usize = 200;
const MCP_RESOURCE_DESCRIPTION_LIMIT: usize = KIB;
const WEB_FETCH_LIMIT: usize = 32 * KIB;
const WEB_SEARCH_LIMIT: usize = 16 * KIB;
const WEB_SEARCH_TITLE_LIMIT: usize = 512;
const WEB_SEARCH_SNIPPET_LIMIT: usize = 2 * 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 tool_completion(
tool_name: &str,
tool: &mut pb::tool_call::Tool,
content: &mut String,
) {
use pb::tool_call::Tool;
match tool {
Tool::ShellToolCall(tool) => gate_shell(tool),
Tool::GrepToolCall(tool) => gate_grep(tool),
Tool::GlobToolCall(tool) => gate_glob(tool),
Tool::ReadToolCall(tool) => gate_read(tool),
Tool::EditToolCall(tool) => gate_edit(tool_name, tool),
Tool::McpToolCall(tool) => gate_mcp(tool),
Tool::ListMcpResourcesToolCall(tool) => gate_mcp_resources(tool),
Tool::ReadMcpResourceToolCall(tool) => gate_mcp_resource(tool),
Tool::GetMcpToolsToolCall(tool) => gate_mcp_tools(tool),
Tool::WebFetchToolCall(tool) => gate_web_fetch(tool),
Tool::WebSearchToolCall(tool) => gate_web_search(tool),
Tool::GenerateImageToolCall(tool) => gate_generate_image(tool),
_ => {}
}
*content = limit_tool_result_text(tool_name, content);
}
pub(super) fn exec_message(message: &mut pb::exec_client_message::Message) {
@@ -20,6 +61,12 @@ pub(super) fn exec_message(message: &mut pb::exec_client_message::Message) {
}
}
fn gate_shell(tool: &mut pb::ShellToolCall) {
if let Some(result) = tool.result.as_mut() {
gate_shell_result(result);
}
}
fn gate_shell_result(result: &mut pb::ShellResult) {
use pb::shell_result::Result;
match result.result.as_mut() {
@@ -27,22 +74,584 @@ fn gate_shell_result(result: &mut pb::ShellResult) {
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);
*interleaved = truncate_edges(
"Shell interleaved output",
interleaved,
SHELL_INTERLEAVED_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);
*interleaved = truncate_edges(
"Shell interleaved output",
interleaved,
SHELL_INTERLEAVED_LIMIT,
);
}
}
_ => {}
}
}
fn gate_read(tool: &mut pb::ReadToolCall) {
let Some(pb::read_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let Some(output) = success.output.as_mut() else {
return;
};
match output {
pb::read_tool_success::Output::Content(value) => {
let next = truncate_text("Read", value, READ_CONTENT_LIMIT);
if next != *value {
*value = next;
success.exceeded_limit = true;
}
}
pb::read_tool_success::Output::Data(value) if value.len() > READ_BINARY_LIMIT => {
let notice = truncation_notice("Read binary data", READ_BINARY_LIMIT, 0, value.len());
success.output = Some(pb::read_tool_success::Output::Content(notice));
success.exceeded_limit = true;
}
_ => {}
}
}
fn gate_glob(tool: &mut pb::GlobToolCall) {
let Some(pb::glob_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let original = success.files.len();
if original <= GLOB_FILE_LIMIT {
if success.total_files <= 0 {
success.total_files = original as i32;
}
return;
}
success.files.truncate(GLOB_FILE_LIMIT);
success.total_files = success.total_files.max(original as i32);
success.client_truncated = true;
}
fn gate_grep(tool: &mut pb::GrepToolCall) {
let Some(pb::grep_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let mut budget = GrepBudget {
content_bytes: GREP_CONTENT_LIMIT,
matches: GREP_TOTAL_MATCHES,
};
let mut workspace_names = success
.workspace_results
.keys()
.cloned()
.collect::<Vec<_>>();
workspace_names.sort_unstable();
for name in workspace_names {
if let Some(result) = success.workspace_results.get_mut(&name) {
gate_grep_union(result, &mut budget);
}
}
if let Some(result) = success.active_editor_result.as_mut() {
gate_grep_union(result, &mut budget);
}
}
struct GrepBudget {
content_bytes: usize,
matches: usize,
}
fn gate_grep_union(result: &mut pb::GrepUnionResult, budget: &mut GrepBudget) {
use pb::grep_union_result::Result;
match result.result.as_mut() {
Some(Result::Content(content)) => gate_grep_content(content, budget),
Some(Result::Files(files)) => {
let original = files.files.len();
if original > GREP_LIST_LIMIT {
files.files.truncate(GREP_LIST_LIMIT);
files.client_truncated = true;
}
if files.total_files <= 0 {
files.total_files = original as i32;
}
}
Some(Result::Count(counts)) => {
let original = counts.counts.len();
if original > GREP_LIST_LIMIT {
counts.counts.truncate(GREP_LIST_LIMIT);
counts.client_truncated = true;
}
if counts.total_files <= 0 {
counts.total_files = original as i32;
}
}
None => {}
}
}
fn gate_grep_content(content: &mut pb::GrepContentResult, budget: &mut GrepBudget) {
if content
.matches
.iter()
.flat_map(|file| &file.matches)
.any(is_grep_notice)
{
return;
}
let original_bytes = grep_content_bytes(&content.matches);
let original_files = content.matches.len();
let mut truncated = false;
let mut files = Vec::with_capacity(original_files);
for file in &content.matches {
if budget.matches == 0 || budget.content_bytes == 0 {
truncated = true;
break;
}
let mut next = pb::GrepFileMatch {
file: file.file.clone(),
matches: Vec::new(),
};
for matched in &file.matches {
if is_grep_notice(matched) {
next.matches.push(matched.clone());
continue;
}
if next.matches.len() >= GREP_MATCHES_PER_FILE
|| budget.matches == 0
|| budget.content_bytes == 0
{
truncated = true;
break;
}
let mut next_match = matched.clone();
let original = next_match.content.clone();
next_match.content = truncate_text("Grep match", &original, GREP_MATCH_LIMIT);
if next_match.content != original {
next_match.content_truncated = true;
truncated = true;
}
if next_match.content.len() > budget.content_bytes {
next_match.content =
truncate_text("Grep", &next_match.content, budget.content_bytes);
next_match.content_truncated = true;
truncated = true;
}
if next_match.content.trim().is_empty() {
truncated = true;
break;
}
budget.content_bytes -= next_match.content.len();
budget.matches -= 1;
next.matches.push(next_match);
}
if next.matches.len() < file.matches.len() {
truncated = true;
}
if !next.matches.is_empty() {
files.push(next);
}
}
if files.len() < original_files {
truncated = true;
}
if truncated {
content.client_truncated = true;
add_grep_notice(&mut files, original_bytes);
}
content.matches = files;
}
fn add_grep_notice(files: &mut Vec<pb::GrepFileMatch>, original_bytes: usize) {
if files
.iter()
.flat_map(|file| &file.matches)
.any(is_grep_notice)
{
return;
}
loop {
let used = grep_content_bytes(files);
let notice = truncation_notice("Grep", GREP_CONTENT_LIMIT, used, original_bytes);
if used.saturating_add(notice.len()) <= GREP_CONTENT_LIMIT {
let matched = pb::GrepContentMatch {
line_number: 0,
content: notice,
content_truncated: true,
is_context_line: true,
};
if let Some(file) = files.last_mut() {
file.matches.push(matched);
} else {
files.push(pb::GrepFileMatch {
file: "[truncated]".into(),
matches: vec![matched],
});
}
return;
}
let Some(file) = files.last_mut() else {
return;
};
file.matches.pop();
if file.matches.is_empty() {
files.pop();
}
}
}
fn is_grep_notice(matched: &pb::GrepContentMatch) -> bool {
matched.line_number == 0
&& matched.content_truncated
&& matched
.content
.starts_with("[truncated: Grep result exceeded")
}
fn grep_content_bytes(files: &[pb::GrepFileMatch]) -> usize {
files
.iter()
.flat_map(|file| &file.matches)
.map(|matched| matched.content.len())
.sum()
}
fn gate_edit(tool_name: &str, tool: &mut pb::EditToolCall) {
let Some(pb::edit_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let limit = match tool_name.trim() {
"PatchEdit" | "PatchEditLines" | "PatchEditSpan" | "StrReplace" => PATCH_EDIT_RESULT_LIMIT,
_ => EDIT_RESULT_LIMIT,
};
if let Some(diff) = success.diff_string.as_mut() {
*diff = truncate_text(tool_name, diff, limit);
success.before_full_file_content = None;
success.after_full_file_content.clear();
} else {
success.before_full_file_content = None;
success.after_full_file_content =
truncate_text(tool_name, &success.after_full_file_content, limit);
}
}
fn gate_mcp(tool: &mut pb::McpToolCall) {
let Some(pb::mcp_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if success.content.iter().any(is_mcp_notice) {
return;
}
let mut notices = Vec::new();
if structured_json_len(&success.structured_content) > MCP_STRUCTURED_LIMIT {
let original = structured_json_len(&success.structured_content);
success.structured_content = truncated_struct(original, MCP_STRUCTURED_LIMIT);
notices.push(truncation_notice(
"MCP structured_content",
MCP_STRUCTURED_LIMIT,
0,
original,
));
}
let original_items = success.content.len();
if original_items > MCP_CONTENT_ITEM_LIMIT {
success.content.truncate(MCP_CONTENT_ITEM_LIMIT);
notices.push(format!(
"[truncated: MCP content items exceeded {MCP_CONTENT_ITEM_LIMIT} items; showing {MCP_CONTENT_ITEM_LIMIT} of {original_items} items]"
));
}
let mut remaining_text = MCP_TEXT_LIMIT;
let mut content = Vec::with_capacity(success.content.len() + notices.len());
for mut item in std::mem::take(&mut success.content) {
match item.content.as_mut() {
Some(pb::mcp_tool_result_content_item::Content::Text(text)) => {
let original = text.text.clone();
let next = truncate_text("MCP content item", &original, MCP_TEXT_LIMIT);
if remaining_text == 0 {
notices.push(truncation_notice(
"MCP text",
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT.saturating_add(original.len()),
));
continue;
}
text.text = truncate_text("MCP text", &next, remaining_text);
remaining_text = remaining_text.saturating_sub(text.text.len());
}
Some(pb::mcp_tool_result_content_item::Content::Image(image))
if image.data.len() > MCP_BINARY_LIMIT =>
{
let original = image.data.len();
image.data.truncate(MCP_BINARY_LIMIT);
notices.push(truncation_notice(
"MCP image data",
MCP_BINARY_LIMIT,
image.data.len(),
original,
));
}
_ => {}
}
content.push(item);
}
content.extend(notices.into_iter().map(mcp_notice));
success.content = content;
}
fn mcp_notice(text: String) -> pb::McpToolResultContentItem {
pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text,
output_location: None,
},
)),
}
}
fn is_mcp_notice(item: &pb::McpToolResultContentItem) -> bool {
matches!(
item.content.as_ref(),
Some(pb::mcp_tool_result_content_item::Content::Text(text))
if text.text.starts_with("[truncated:")
)
}
fn structured_json_len(value: &Option<prost_types::Struct>) -> usize {
value
.as_ref()
.and_then(|value| {
serde_json::to_vec(&serde_json::Value::Object(
value
.fields
.iter()
.map(|(key, value)| (key.clone(), super::prost_json(value)))
.collect(),
))
.ok()
})
.map_or(0, |value| value.len())
}
fn truncated_struct(original: usize, limit: usize) -> Option<prost_types::Struct> {
Some(prost_types::Struct {
fields: BTreeMap::from([
("_truncated".into(), prost_bool(true)),
("original_json_bytes".into(), prost_number(original as f64)),
("limit_bytes".into(), prost_number(limit as f64)),
]),
})
}
fn prost_bool(value: bool) -> prost_types::Value {
prost_types::Value {
kind: Some(prost_types::value::Kind::BoolValue(value)),
}
}
fn prost_number(value: f64) -> prost_types::Value {
prost_types::Value {
kind: Some(prost_types::value::Kind::NumberValue(value)),
}
}
fn gate_mcp_resources(tool: &mut pb::ListMcpResourcesToolCall) {
let Some(pb::list_mcp_resources_exec_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if success
.resources
.iter()
.any(|resource| resource.uri == "truncated:list-mcp-resources")
{
return;
}
let original = success.resources.len();
success.resources.truncate(MCP_RESOURCE_LIMIT);
for resource in &mut success.resources {
if let Some(description) = resource.description.as_mut() {
*description = truncate_text(
"MCP resource description",
description,
MCP_RESOURCE_DESCRIPTION_LIMIT,
);
}
}
if success.resources.len() < original {
success
.resources
.push(pb::list_mcp_resources_exec_result::McpResource {
uri: "truncated:list-mcp-resources".into(),
name: Some("truncated".into()),
description: Some(truncation_notice(
"ListMcpResources",
MCP_TEXT_LIMIT,
success.resources.len(),
original,
)),
..Default::default()
});
}
}
fn gate_mcp_resource(tool: &mut pb::ReadMcpResourceToolCall) {
let Some(pb::read_mcp_resource_exec_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
match success.content.as_mut() {
Some(pb::read_mcp_resource_success::Content::Text(text)) => {
*text = truncate_text("FetchMcpResource", text, MCP_TEXT_LIMIT);
}
Some(pb::read_mcp_resource_success::Content::Blob(blob))
if blob.len() > MCP_BINARY_LIMIT =>
{
let notice =
truncation_notice("FetchMcpResource blob", MCP_BINARY_LIMIT, 0, blob.len());
success.content = Some(pb::read_mcp_resource_success::Content::Text(notice));
}
_ => {}
}
}
fn gate_mcp_tools(tool: &mut pb::GetMcpToolsToolCall) {
let Some(pb::get_mcp_tools_agent_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
success.content = truncate_text("GetMcpTools", &success.content, MCP_TEXT_LIMIT);
}
fn gate_web_fetch(tool: &mut pb::WebFetchToolCall) {
let Some(pb::web_fetch_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
success.markdown = truncate_text("WebFetch", &success.markdown, WEB_FETCH_LIMIT);
}
fn gate_web_search(tool: &mut pb::WebSearchToolCall) {
let Some(pb::web_search_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
for reference in &mut success.references {
reference.title =
truncate_text("WebSearch title", &reference.title, WEB_SEARCH_TITLE_LIMIT);
reference.chunk = truncate_text(
"WebSearch snippet",
&reference.chunk,
WEB_SEARCH_SNIPPET_LIMIT,
);
}
let original = web_search_bytes(&success.references);
while success.references.len() > 1 && web_search_bytes(&success.references) > WEB_SEARCH_LIMIT {
success.references.pop();
}
if original > WEB_SEARCH_LIMIT {
let total = web_search_bytes(&success.references);
if let Some(reference) = success.references.last_mut() {
let other = total.saturating_sub(reference.chunk.len());
let notice = truncation_notice(
"WebSearch",
WEB_SEARCH_LIMIT,
WEB_SEARCH_LIMIT.saturating_sub(other),
original,
);
let available = WEB_SEARCH_LIMIT.saturating_sub(other + notice.len() + 2);
reference.chunk = format!(
"{}\n\n{notice}",
utf8_prefix(&reference.chunk, available).trim_end_matches('\n')
);
}
}
}
fn web_search_bytes(references: &[pb::WebSearchReference]) -> usize {
references
.iter()
.map(|reference| reference.title.len() + reference.url.len() + reference.chunk.len())
.sum()
}
fn gate_generate_image(tool: &mut pb::GenerateImageToolCall) {
let Some(pb::generate_image_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if !success.image_data.trim().is_empty()
&& !success
.image_data
.starts_with("[base64 image data omitted from replay; bytes=")
{
let original = success.image_data.trim().len();
success.image_data = format!("[base64 image data omitted from replay; bytes={original}]");
}
}
fn truncate_text(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; showing {shown} of {original} bytes]"
);
let available = limit.saturating_sub(notice.len());
let kept = utf8_prefix(content, available);
if kept.len() == shown {
return format!("{}{notice}", kept.trim_end_matches('\n'));
}
shown = kept.len();
}
}
fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
if content.len() <= limit {
return content.to_string();
@@ -64,6 +673,12 @@ fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
}
}
fn truncation_notice(tool_name: &str, limit: usize, shown: usize, original: usize) -> String {
format!(
"[truncated: {tool_name} result exceeded {limit} bytes; showing {shown} of {original} bytes]"
)
}
fn utf8_prefix(value: &str, limit: usize) -> &str {
let mut end = limit.min(value.len());
while end > 0 && !value.is_char_boundary(end) {
@@ -84,33 +699,212 @@ fn utf8_suffix(value: &str, limit: usize) -> &str {
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));
let mut tool = pb::tool_call::Tool::ShellToolCall(pb::ShellToolCall::default());
model_content(&shell_tool(), &mut content);
tool_completion("Shell", &mut tool, &mut content);
assert!(content.len() <= SHELL_CONTENT_LIMIT);
assert!(content.len() <= 128 * KIB);
assert!(content.starts_with("HEAD"));
assert!(content.ends_with("TAIL"));
assert!(content.contains("omitted middle"));
assert!(content.contains("[truncated: Shell result exceeded"));
}
#[test]
fn non_shell_output_is_unchanged() {
fn grep_limits_matches_per_file_total_bytes_and_adds_notice() {
let matches = (0..150)
.map(|line_number| pb::GrepContentMatch {
line_number,
content: "x".repeat(3 * KIB),
..Default::default()
})
.collect();
let mut tool = pb::tool_call::Tool::GrepToolCall(pb::GrepToolCall {
result: Some(pb::GrepResult {
result: Some(pb::grep_result::Result::Success(pb::GrepSuccess {
workspace_results: std::collections::HashMap::from([(
"workspace".into(),
pb::GrepUnionResult {
result: Some(pb::grep_union_result::Result::Content(
pb::GrepContentResult {
matches: vec![pb::GrepFileMatch {
file: "large.txt".into(),
matches,
}],
..Default::default()
},
)),
},
)]),
..Default::default()
})),
}),
..Default::default()
});
let mut model_content = "x".repeat(128 * KIB);
tool_completion("Grep", &mut tool, &mut model_content);
assert!(model_content.len() <= GREP_CONTENT_LIMIT);
assert!(model_content.contains("[truncated: Grep result exceeded"));
let pb::tool_call::Tool::GrepToolCall(tool) = tool else {
unreachable!()
};
let Some(pb::grep_result::Result::Success(success)) =
tool.result.clone().and_then(|result| result.result)
else {
panic!("expected grep success")
};
let result = success.workspace_results.get("workspace").unwrap();
let Some(pb::grep_union_result::Result::Content(content)) = result.result.as_ref() else {
panic!("expected grep content")
};
assert!(content.client_truncated);
assert!(grep_content_bytes(&content.matches) <= GREP_CONTENT_LIMIT);
assert!(content.matches[0].matches.len() <= GREP_MATCHES_PER_FILE + 1);
assert!(content.matches[0]
.matches
.last()
.unwrap()
.content
.contains("[truncated: Grep result exceeded"));
let once = tool.clone();
let mut tool_enum = pb::tool_call::Tool::GrepToolCall(tool);
let mut second_content = model_content.clone();
tool_completion("Grep", &mut tool_enum, &mut second_content);
let pb::tool_call::Tool::GrepToolCall(second) = tool_enum else {
panic!("expected grep tool")
};
assert_eq!(second, once);
assert_eq!(second_content, model_content);
}
#[test]
fn read_content_is_limited_and_marked() {
let mut tool = pb::tool_call::Tool::ReadToolCall(pb::ReadToolCall {
result: Some(pb::ReadToolResult {
result: Some(pb::read_tool_result::Result::Success(pb::ReadToolSuccess {
output: Some(pb::read_tool_success::Output::Content(
"前".repeat(READ_CONTENT_LIMIT),
)),
..Default::default()
})),
}),
..Default::default()
});
let mut content = "前".repeat(READ_CONTENT_LIMIT);
tool_completion("Read", &mut tool, &mut content);
assert!(content.len() <= READ_CONTENT_LIMIT);
let pb::tool_call::Tool::ReadToolCall(tool) = tool else {
unreachable!()
};
let Some(pb::read_tool_result::Result::Success(success)) =
tool.result.and_then(|result| result.result)
else {
panic!("expected read success")
};
assert!(success.exceeded_limit);
let Some(pb::read_tool_success::Output::Content(output)) = success.output else {
panic!("expected text output")
};
assert!(output.len() <= READ_CONTENT_LIMIT);
assert!(output.contains("[truncated: Read result exceeded"));
}
#[test]
fn mcp_limits_items_text_and_structured_content() {
let mut tool = pb::tool_call::Tool::McpToolCall(pb::McpToolCall {
result: Some(pb::McpToolResult {
result: Some(pb::mcp_tool_result::Result::Success(pb::McpSuccess {
content: (0..25)
.map(|_| pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text: "x".repeat(4 * KIB),
..Default::default()
},
)),
})
.collect(),
structured_content: Some(prost_types::Struct {
fields: BTreeMap::from([(
"large".into(),
prost_types::Value {
kind: Some(prost_types::value::Kind::StringValue(
"x".repeat(64 * KIB),
)),
},
)]),
}),
..Default::default()
})),
}),
..Default::default()
});
let mut content = "x".repeat(64 * KIB);
let original = content.clone();
model_content(
&pb::tool_call::Tool::ReadToolCall(pb::ReadToolCall::default()),
&mut content,
tool_completion("CallMcpTool", &mut tool, &mut content);
assert!(content.len() <= MCP_TEXT_LIMIT);
let pb::tool_call::Tool::McpToolCall(tool) = tool else {
unreachable!()
};
let Some(pb::mcp_tool_result::Result::Success(success)) =
tool.result.and_then(|result| result.result)
else {
panic!("expected mcp success")
};
assert!(success.content.len() > MCP_CONTENT_ITEM_LIMIT);
assert_eq!(
success
.structured_content
.unwrap()
.fields
.get("_truncated")
.unwrap()
.kind,
Some(prost_types::value::Kind::BoolValue(true))
);
assert!(success.content.iter().any(|item| matches!(
item.content.as_ref(),
Some(pb::mcp_tool_result_content_item::Content::Text(text))
if text.text.contains("MCP content items exceeded")
)));
}
assert_eq!(content, original);
#[test]
fn edit_keeps_only_a_bounded_diff() {
let mut tool = pb::tool_call::Tool::EditToolCall(pb::EditToolCall {
result: Some(pb::EditResult {
result: Some(pb::edit_result::Result::Success(pb::EditSuccess {
diff_string: Some("d".repeat(16 * KIB)),
before_full_file_content: Some("b".repeat(64 * KIB)),
after_full_file_content: "a".repeat(64 * KIB),
..Default::default()
})),
}),
..Default::default()
});
let mut content = "x".repeat(64 * KIB);
tool_completion("StrReplace", &mut tool, &mut content);
assert!(content.len() <= PATCH_EDIT_RESULT_LIMIT);
let pb::tool_call::Tool::EditToolCall(tool) = tool else {
unreachable!()
};
let Some(pb::edit_result::Result::Success(success)) =
tool.result.and_then(|result| result.result)
else {
panic!("expected edit success")
};
assert!(success.diff_string.unwrap().len() <= PATCH_EDIT_RESULT_LIMIT);
assert!(success.before_full_file_content.is_none());
assert!(success.after_full_file_content.is_empty());
}
#[test]
@@ -139,31 +933,6 @@ mod tests {
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);
assert!(success.interleaved_output.unwrap().len() <= SHELL_INTERLEAVED_LIMIT);
}
}
+5 -4
View File
@@ -1,4 +1,3 @@
mod await_shell;
mod exec;
mod gate;
mod interaction;
@@ -19,7 +18,6 @@ use crate::{
use super::runtime::now_ms;
pub(crate) use await_shell::{await_error, await_result, await_sleep};
pub(crate) use exec::{edit_failure, from_exec};
pub(crate) use interaction::{complete_web_fetch, complete_web_search, from_interaction};
pub(crate) use local::{local, subagents_disabled, todo_items};
@@ -89,9 +87,12 @@ impl ToolCompletion {
call: &ToolCall,
started_at_ms: u64,
mut result: ToolResult,
tool: pb::tool_call::Tool,
mut tool: pb::tool_call::Tool,
) -> Self {
gate::model_content(&tool, &mut result.content);
// Apply the model-visible size gate once, at the tool completion
// boundary. Canonical history and every provider projection then
// carry the same bounded result without reprocessing it.
gate::tool_completion(&call.name, &mut tool, &mut result.content);
Self {
result,
tool_call: pb::ToolCall {
+36 -64
View File
@@ -1,10 +1,9 @@
use std::{
collections::HashMap,
collections::{HashMap, HashSet},
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
time::Instant,
};
use tokio::sync::Mutex;
@@ -19,6 +18,7 @@ pub struct CursorToolRuntime {
execs: Arc<Mutex<HashMap<u32, PendingExec>>>,
interactions: Arc<Mutex<HashMap<u32, PendingInteraction>>>,
completed: Arc<Mutex<HashMap<u32, String>>>,
interrupted: Arc<Mutex<HashSet<u32>>>,
}
pub(crate) struct PendingExec {
@@ -35,14 +35,6 @@ pub(crate) enum ExecStage {
DynamicMcp(pb::McpToolDefinition),
EditRead,
EditWrite(EditWrite),
Await(AwaitState),
}
pub(crate) struct AwaitState {
pub deadline: Instant,
pub output_file_path: String,
pub task_id: String,
pub regex: Option<String>,
}
#[derive(Clone, Debug, Default)]
@@ -171,60 +163,6 @@ impl CursorToolRuntime {
.await
}
pub(crate) async fn reserve_await(
&self,
call: &ToolCall,
context: &ExecContext,
) -> Result<u32> {
let task_id = call
.arguments
.get("shell_id")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol("AwaitShell is missing shell_id".into()))?;
let block_ms = call
.arguments
.get("block_until_ms")
.and_then(serde_json::Value::as_u64)
.unwrap_or(30_000);
if block_ms > 7_140_000 {
return Err(Error::Protocol(
"AwaitShell block_until_ms exceeds 7140000".into(),
));
}
let output_file_path = format!(
"{}/{}.txt",
context.terminals_folder.trim_end_matches('/'),
task_id
);
self.reserve_exec_stage(
call,
context,
ExecStage::Await(AwaitState {
deadline: Instant::now() + std::time::Duration::from_millis(block_ms),
output_file_path,
task_id: task_id.to_string(),
regex: call
.arguments
.get("pattern")
.and_then(serde_json::Value::as_str)
.map(str::to_string),
}),
None,
)
.await
}
pub(crate) async fn reserve_await_again(
&self,
call: &ToolCall,
context: &ExecContext,
state: AwaitState,
started_at_ms: u64,
) -> Result<u32> {
self.reserve_exec_stage(call, context, ExecStage::Await(state), Some(started_at_ms))
.await
}
async fn reserve_exec_stage(
&self,
call: &ToolCall,
@@ -311,6 +249,10 @@ impl CursorToolRuntime {
self.completed.lock().await.get(&id).cloned()
}
pub async fn is_interrupted(&self, id: u32) -> bool {
self.interrupted.lock().await.contains(&id)
}
pub async fn clear_completed(&self) {
self.completed.lock().await.clear();
}
@@ -329,9 +271,39 @@ impl CursorToolRuntime {
ids.sort_unstable();
self.interactions.lock().await.clear();
self.completed.lock().await.clear();
self.interrupted.lock().await.clear();
ids
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
let (abort_ids, interrupted_ids) = {
let mut entries = self.execs.lock().await;
let mut abort_ids = Vec::new();
let mut interrupted_ids = Vec::new();
entries.retain(|id, entry| {
interrupted_ids.push(*id);
let keep_running = entry.call.name.eq_ignore_ascii_case("Task");
if !keep_running {
abort_ids.push(*id);
}
keep_running
});
(abort_ids, interrupted_ids)
};
let interaction_ids = {
let mut interactions = self.interactions.lock().await;
let ids = interactions.keys().copied().collect::<Vec<_>>();
interactions.clear();
ids
};
let mut interrupted = self.interrupted.lock().await;
interrupted.extend(interrupted_ids);
interrupted.extend(interaction_ids);
let mut abort_ids = abort_ids;
abort_ids.sort_unstable();
abort_ids
}
pub async fn running_exec_ids(&self) -> Vec<u32> {
let mut ids = self.execs.lock().await.keys().copied().collect::<Vec<_>>();
ids.sort_unstable();
+5
View File
@@ -23,6 +23,11 @@ pub(super) struct DeferredEdit {
}
impl EditSchedule {
pub fn clear(&mut self) {
self.paths.clear();
self.active_paths.clear();
}
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);
+78
View File
@@ -181,6 +181,11 @@ impl ModelConfig {
pub fn configure(&self, model: &mut super::ModelSpec) {
model.display_name = Some(self.display_name.clone());
// A request-selected context is authoritative. Use the saved model
// value only when Cursor did not send a context parameter.
if model.context_window_tokens.is_none() {
model.context_window_tokens = self.context_window_tokens;
}
if model.reasoning.effort.is_none() {
model.reasoning.effort = match self.model_type {
ModelType::OpenAi => self.reasoning_effort.clone(),
@@ -505,4 +510,77 @@ mod tests {
"https://example.com/v1/messages"
);
}
#[test]
fn configured_context_window_does_not_override_the_client_request() {
let input = input();
let config = ModelConfig {
model_hash: "hash".into(),
sort_order: input.sort_order,
display_name: input.display_name,
model_type: input.model_type,
base_url: input.base_url,
use_full_url: input.use_full_url,
api_key: input.api_key,
tooltip_data: input.tooltip_data,
model_id: input.model_id,
reasoning_effort: input.reasoning_effort,
openai_endpoint: input.openai_endpoint,
openai_extra_params_enabled: input.openai_extra_params_enabled,
openai_extra_params: input.openai_extra_params,
custom_headers_enabled: input.custom_headers_enabled,
custom_headers: input.custom_headers,
anthropic_extra_params_enabled: input.anthropic_extra_params_enabled,
anthropic_extra_params: input.anthropic_extra_params,
context_window_tokens: Some(350_000),
max_completion_tokens: input.max_completion_tokens,
anthropic_max_tokens: input.anthropic_max_tokens,
anthropic_thinking_effort: input.anthropic_thinking_effort,
thinking_budget_tokens: input.thinking_budget_tokens,
created_at_ms: 0,
updated_at_ms: 0,
};
let mut requested = super::super::ModelSpec::new("model-a");
requested.context_window_tokens = Some(200_000);
config.configure(&mut requested);
assert_eq!(requested.context_window_tokens, Some(200_000));
}
#[test]
fn configured_context_window_fills_missing_client_value() {
let input = input();
let config = ModelConfig {
model_hash: "hash".into(),
sort_order: input.sort_order,
display_name: input.display_name,
model_type: input.model_type,
base_url: input.base_url,
use_full_url: input.use_full_url,
api_key: input.api_key,
tooltip_data: input.tooltip_data,
model_id: input.model_id,
reasoning_effort: input.reasoning_effort,
openai_endpoint: input.openai_endpoint,
openai_extra_params_enabled: input.openai_extra_params_enabled,
openai_extra_params: input.openai_extra_params,
custom_headers_enabled: input.custom_headers_enabled,
custom_headers: input.custom_headers,
anthropic_extra_params_enabled: input.anthropic_extra_params_enabled,
anthropic_extra_params: input.anthropic_extra_params,
context_window_tokens: Some(350_000),
max_completion_tokens: input.max_completion_tokens,
anthropic_max_tokens: input.anthropic_max_tokens,
anthropic_thinking_effort: input.anthropic_thinking_effort,
thinking_budget_tokens: input.thinking_budget_tokens,
created_at_ms: 0,
updated_at_ms: 0,
};
let mut requested = super::super::ModelSpec::new("model-a");
config.configure(&mut requested);
assert_eq!(requested.context_window_tokens, Some(350_000));
}
}
+2
View File
@@ -11,6 +11,7 @@ mod run;
mod runtime_tag;
mod token_count;
mod tool;
mod tool_result_replay;
mod usage;
pub use configuration::*;
@@ -26,4 +27,5 @@ pub use run::*;
pub use runtime_tag::*;
pub(crate) use token_count::*;
pub use tool::*;
pub(crate) use tool_result_replay::limit_tool_result_text;
pub use usage::*;
+226
View File
@@ -0,0 +1,226 @@
use serde_json::Value;
const KIB: usize = 1024;
pub(crate) fn limit_tool_result_text(name: &str, content: &str) -> String {
let Some(limit) = replay_limit(name) else {
return content.to_string();
};
let content = match name.trim() {
"GenerateImage" => compact_generate_image(content),
"Shell" => compact_shell(content),
"PatchEdit" | "PatchEditLines" | "PatchEditSpan" | "StrReplace" | "Edit" | "Write" => {
compact_edit(name, content)
}
_ => None,
}
.unwrap_or_else(|| content.to_string());
truncate_replay_text(name, &content, limit)
}
fn replay_limit(name: &str) -> Option<usize> {
match name.trim() {
"GenerateImage" | "WebSearch" => Some(16 * KIB),
"Read" => Some(64 * KIB),
"Shell" => Some(128 * KIB),
"Grep" | "Glob" => Some(32 * KIB),
"PatchEdit" | "PatchEditLines" | "PatchEditSpan" | "StrReplace" => Some(4 * KIB),
"Edit" | "EditNotebook" | "Write" | "WebFetch" => Some(32 * KIB),
"CallMcpTool" | "FetchMcpResource" | "ListMcpResources" | "GetMcpTools"
| "SembleSearch" | "SembleFindRelated" => Some(32 * KIB),
_ => None,
}
}
fn truncate_replay_text(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: {name} result exceeded {limit} bytes; showing {shown} of {original} bytes]"
);
let available = limit.saturating_sub(notice.len());
let kept = utf8_prefix(content, available);
if kept.len() == shown {
return format!("{}{notice}", kept.trim_end_matches('\n'));
}
shown = kept.len();
}
}
fn compact_generate_image(content: &str) -> Option<String> {
let mut value = serde_json::from_str::<Value>(content.trim()).ok()?;
if !replace_image_data(&mut value) {
return None;
}
serde_json::to_string(&value).ok()
}
fn replace_image_data(value: &mut Value) -> bool {
match value {
Value::Object(object) => {
let mut changed = false;
for (key, child) in object.iter_mut() {
if matches!(key.as_str(), "image_data" | "imageData") {
if let Value::String(data) = child {
if data.starts_with("[base64 image data omitted from replay; bytes=") {
continue;
}
*child = Value::String(format!(
"[base64 image data omitted from replay; bytes={}]",
data.trim().len()
));
changed = true;
continue;
}
}
changed |= replace_image_data(child);
}
changed
}
Value::Array(items) => items.iter_mut().any(replace_image_data),
_ => false,
}
}
fn compact_shell(content: &str) -> Option<String> {
let mut value = serde_json::from_str::<Value>(content.trim()).ok()?;
if !compact_shell_fields(&mut value) {
return None;
}
serde_json::to_string(&value).ok()
}
fn compact_shell_fields(value: &mut Value) -> bool {
match value {
Value::Object(object) => {
let mut changed = false;
for (key, child) in object.iter_mut() {
if let Value::String(text) = child {
let limit = match key.as_str() {
"stdout" | "stderr" => Some(16 * KIB),
"interleaved_output" | "interleavedOutput" => Some(32 * KIB),
_ => None,
};
if let Some(limit) = limit {
let next = truncate_middle(&format!("Shell {key}"), text, limit);
if next != *text {
*text = next;
changed = true;
}
continue;
}
}
changed |= compact_shell_fields(child);
}
changed
}
Value::Array(items) => items.iter_mut().any(compact_shell_fields),
_ => false,
}
}
fn compact_edit(name: &str, content: &str) -> Option<String> {
let value = serde_json::from_str::<Value>(content.trim()).ok()?;
let success = value.get("success")?.as_object()?;
let diff = success
.get("diff_string")
.or_else(|| success.get("diffString"))
.and_then(Value::as_str)
.filter(|text| !text.is_empty())
.map(|text| truncate_replay_text(name, text, edit_limit(name)));
if let Some(diff) = diff {
return Some(serde_json::json!({"success": {"diff_string": diff}}).to_string());
}
let after = success
.get("after_full_file_content")
.or_else(|| success.get("afterFullFileContent"))
.and_then(Value::as_str)
.filter(|text| !text.is_empty())
.map(|text| truncate_replay_text(name, text, edit_limit(name)));
after
.map(|after| serde_json::json!({"success": {"after_full_file_content": after}}).to_string())
}
fn edit_limit(name: &str) -> usize {
match name.trim() {
"PatchEdit" | "PatchEditLines" | "PatchEditSpan" | "StrReplace" => 4 * KIB,
_ => 32 * KIB,
}
}
fn truncate_middle(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: {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() + tail.len();
let next_notice = format!(
"\n\n[truncated: {name} result exceeded {limit} bytes; omitted middle; showing {next_shown} of {original} bytes]\n\n"
);
let output = format!("{head}{next_notice}{tail}");
if output.len() <= limit || next_notice == notice {
return output;
}
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::*;
#[test]
fn truncation_preserves_utf8_and_limit() {
let content = "前".repeat(32 * KIB);
let truncated = truncate_replay_text("Grep", &content, 32 * KIB);
assert!(truncated.len() <= 32 * KIB);
assert!(truncated.is_char_boundary(truncated.len()));
assert!(truncated.contains("[truncated: Grep result exceeded"));
}
#[test]
fn json_replay_compacts_nested_image_data_and_shell_streams() {
let image = serde_json::json!({"success": {"image_data": "x".repeat(64 * KIB)}});
let image_result = limit_tool_result_text("GenerateImage", &image.to_string());
assert!(image_result.contains("base64 image data omitted"));
assert!(image_result.len() < 1024);
assert_eq!(
limit_tool_result_text("GenerateImage", &image_result),
image_result
);
let shell = serde_json::json!({"success": {"stdout": "x".repeat(64 * KIB)}});
let shell_result = limit_tool_result_text("Shell", &shell.to_string());
assert!(shell_result.len() <= 128 * KIB);
assert!(shell_result.contains("omitted middle"));
}
}
+65 -4
View File
@@ -132,7 +132,7 @@ impl Provider for OpenAiChatProvider {
}
let Some(choice) = value.get("choices").and_then(Value::as_array).and_then(|values| values.first()) else { continue; };
let delta = choice.get("delta").unwrap_or(&Value::Null);
if let Some(reasoning_delta) = delta.get("reasoning_content").and_then(Value::as_str).filter(|text| !text.is_empty()) {
if let Some(reasoning_delta) = delta.get("reasoning_content").or_else(|| delta.get("reasoning")).and_then(Value::as_str).filter(|text| !text.is_empty()) {
if !thinking_open { thinking_open = true; yield ModelEvent::ThinkingStart; }
reasoning.push_str(reasoning_delta);
yield ModelEvent::ThinkingDelta(reasoning_delta.into());
@@ -230,12 +230,27 @@ fn openai_chat_messages(instructions: &str, messages: &[ProjectedMessage]) -> Re
calls,
..
} => {
value.insert("content".into(), Value::String(text.clone()));
let replay_reasoning = replay_state
.as_ref()
.filter(|state| state.provider_kind == "openai_chat")
.and_then(|state| state.value.get("reasoning_content"))
.and_then(Value::as_str);
.and_then(Value::as_str)
.filter(|reasoning| !reasoning.is_empty());
// Chat Completions rejects an empty assistant content string. Tool-call
// assistant messages use null content, while an assistant with no visible
// content at all does not need to be sent.
if text.is_empty() && calls.is_empty() && replay_reasoning.is_none() {
continue;
}
value.insert(
"content".into(),
if text.is_empty() {
Value::Null
} else {
Value::String(text.clone())
},
);
if let Some(reasoning) = replay_reasoning {
value.insert("reasoning_content".into(), Value::String(reasoning.into()));
}
@@ -398,7 +413,7 @@ mod tests {
model::{ContentPart, ProjectedContent, ProjectedMessage, ToolResultContent},
model::{ProviderReplayState, Role, ToolCallContent},
};
use serde_json::json;
use serde_json::{json, Value};
#[test]
fn chat_replay_state_is_encoded_as_reasoning_content() {
@@ -430,6 +445,52 @@ mod tests {
assert_eq!(messages[0]["tool_calls"][0]["id"], "call-1");
}
#[test]
fn chat_tool_call_assistant_uses_null_content() {
let messages = openai_chat_messages(
"",
&[ProjectedMessage {
message_id: "test".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: String::new(),
thinking: String::new(),
replay_state: None,
calls: vec![ToolCallContent {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
arguments: json!({"path": "README.md"}),
}],
},
}],
)
.unwrap();
assert_eq!(messages[0]["content"], Value::Null);
assert!(messages[0]["tool_calls"].is_array());
}
#[test]
fn chat_contentless_assistant_is_omitted() {
let messages = openai_chat_messages(
"",
&[ProjectedMessage {
message_id: "test".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: String::new(),
thinking: String::new(),
replay_state: None,
calls: vec![],
},
}],
)
.unwrap();
assert!(messages.is_empty());
}
#[test]
fn another_provider_replay_does_not_invent_chat_reasoning_content() {
let messages = openai_chat_messages(
+82 -14
View File
@@ -249,14 +249,14 @@ impl RunEngine {
let mut pending_insertions = Vec::new();
let cycle = loop {
tokio::select! {
result = &mut cycle => break result,
biased;
command = client.commands.recv() => {
let message = match command {
Some(ClientCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
continue;
}
Some(ClientCommand::RuntimeMessage(message)) => message,
Some(ClientCommand::InterruptWithMessage(message)) => message,
Some(ClientCommand::RuntimeEvent(event)) => event.into_message(),
Some(ClientCommand::Cancel) => {
cycle_cancellation.cancel();
@@ -321,7 +321,8 @@ impl RunEngine {
Err(outcome) => return (outcome, usage),
};
continue 'model;
}
},
result = &mut cycle => break result,
}
};
let cycle = match cycle {
@@ -561,23 +562,68 @@ impl RunEngine {
let cycle_cancellation = cancellation.child_token();
let (silent_events, mut discarded_events) = tokio::sync::mpsc::channel(256);
let drain = tokio::spawn(async move { while discarded_events.recv().await.is_some() {} });
let cycle = consume_model_cycle(
self.provider.stream(invocation, cycle_cancellation.clone()),
&silent_events,
&cycle_cancellation,
)
.await;
let mut pending_insertions = Vec::new();
let mut interrupted_message = None;
let cycle = {
let cycle = consume_model_cycle(
self.provider.stream(invocation, cycle_cancellation.clone()),
&silent_events,
&cycle_cancellation,
);
tokio::pin!(cycle);
loop {
tokio::select! {
biased;
command = client.commands.recv() => match command {
Some(ClientCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
}
Some(ClientCommand::InterruptWithMessage(message)) => {
cycle_cancellation.cancel();
interrupted_message = Some(message);
break cycle.await;
}
Some(ClientCommand::RuntimeEvent(event)) => {
cycle_cancellation.cancel();
interrupted_message = Some(event.into_message());
break cycle.await;
}
Some(ClientCommand::Cancel) => {
cycle_cancellation.cancel();
return Err(RunOutcome::Cancelled);
}
Some(ClientCommand::ClientClosed { error }) => {
cycle_cancellation.cancel();
return Err(RunOutcome::Failed(RunFailure::Client(error)));
}
Some(ClientCommand::ToolResult(_)) => {
cycle_cancellation.cancel();
return Err(RunOutcome::Failed(RunFailure::Protocol(
"received a tool result while automatic compaction was running".into(),
)));
}
None => {
cycle_cancellation.cancel();
return Err(client_failure());
}
},
result = &mut cycle => break result,
}
}
};
drop(silent_events);
let _ = drain.await;
let (summary, compaction_usage) = match cycle {
Ok(cycle) if cycle.calls.is_empty() && !cycle.text.trim().is_empty() => {
let (summary, compaction_usage) = match (interrupted_message.is_some(), cycle) {
(true, Ok(cycle)) => (fallback_summary(&compactable), cycle.usage),
(true, Err(failure)) => (fallback_summary(&compactable), failure.usage),
(false, Ok(cycle)) if cycle.calls.is_empty() && !cycle.text.trim().is_empty() => {
(cycle.text.trim().to_string(), cycle.usage)
}
Ok(cycle) => {
(false, Ok(cycle)) => {
tracing::warn!("automatic compaction returned no usable summary; using fallback");
(fallback_summary(&compactable), cycle.usage)
}
Err(failure) => {
(false, Err(failure)) => {
tracing::warn!(error = ?failure.failure, "automatic compaction model failed; using fallback");
(fallback_summary(&compactable), failure.usage)
}
@@ -597,7 +643,7 @@ impl RunEngine {
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
replacement.push(summary_message);
replacement.extend(prepared.initial_messages.iter().cloned());
let revision = self
let mut revision = self
.store
.replace_revision(
&prepared.conversation_id,
@@ -623,6 +669,28 @@ impl RunEngine {
emit(client, ClientEvent::AutoCompactionCompleted)
.await
.map_err(|_| client_failure())?;
revision = append_insertions(
&self.store,
prepared,
client,
cancellation,
revision,
pending_insertions,
)
.await?
.0;
if let Some(message) = interrupted_message {
revision = append_runtime_message(
&self.store,
prepared,
client,
cancellation,
revision,
message,
)
.await?
.0;
}
Ok((revision, compaction_usage))
}
}
+102 -30
View File
@@ -1,3 +1,5 @@
use std::collections::HashSet;
use tokio_util::sync::CancellationToken;
use crate::{
@@ -5,7 +7,7 @@ use crate::{
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
StateCommitted,
},
model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId},
model::{PreparedRun, RevisionId, ToolCall, ToolResult, ToolRoundAssistant, ToolRoundId},
store::Store,
};
@@ -70,6 +72,7 @@ pub(super) async fn execute(
.await?;
let mut remaining = calls.len();
let mut completed_call_ids = HashSet::new();
let mut pending_runtime_messages = insertions
.into_iter()
.map(PendingRuntimeMessage::Insertion)
@@ -92,6 +95,7 @@ pub(super) async fn execute(
.await
.map_err(failed)?;
revision = committed.revision_id;
completed_call_ids.insert(call_id.clone());
tracing::info!(
round_id = %round_id,
call_id,
@@ -113,7 +117,10 @@ pub(super) async fn execute(
ClientEvent::StateCommitted(StateCommitted {
revision_id: revision,
tool_round_version: committed.tool_round_version,
cause: CommitCause::ToolResult { call_id },
cause: CommitCause::ToolResult {
call_id,
interrupted: false,
},
barrier,
}),
)
@@ -125,8 +132,66 @@ pub(super) async fn execute(
Some(ClientCommand::RuntimeEvent(event)) => {
pending_runtime_messages.push(PendingRuntimeMessage::Message(event.into_message()));
}
Some(ClientCommand::RuntimeMessage(message)) => {
pending_runtime_messages.push(PendingRuntimeMessage::Message(message));
Some(ClientCommand::InterruptWithMessage(message)) => {
for call in calls
.iter()
.filter(|call| !completed_call_ids.contains(&call.call_id))
{
let result = ToolResult {
call_id: call.call_id.clone(),
content: "Tool execution was interrupted by a newer user message.".into(),
is_error: true,
image: None,
};
let committed = store
.commit_tool_result(
&prepared.conversation_id,
&prepared.run_id,
&round_id,
&result,
)
.await
.map_err(failed)?;
revision = committed.revision_id;
let (barrier, ready) = if committed.settled {
let (barrier, ready) = CommitBarrier::before_continue();
(barrier, Some(ready))
} else {
(CommitBarrier::None, None)
};
send(
client,
ClientEvent::StateCommitted(StateCommitted {
revision_id: revision,
tool_round_version: committed.tool_round_version,
cause: CommitCause::ToolResult {
call_id: call.call_id.clone(),
interrupted: true,
},
barrier,
}),
)
.await?;
if let Some(ready) = ready {
super::engine::wait_for_state_ready(ready, cancellation).await?;
}
}
for pending in pending_runtime_messages {
revision =
append_pending(store, prepared, client, cancellation, revision, pending)
.await?;
}
revision = super::engine::append_runtime_message(
store,
prepared,
client,
cancellation,
revision,
message,
)
.await?
.0;
return Ok(revision);
}
Some(ClientCommand::InsertMessages(insertion)) => {
pending_runtime_messages.push(PendingRuntimeMessage::Insertion(insertion))
@@ -139,32 +204,7 @@ pub(super) async fn execute(
}
}
for pending in pending_runtime_messages {
match pending {
PendingRuntimeMessage::Message(message) => {
revision = super::engine::append_runtime_message(
store,
prepared,
client,
cancellation,
revision,
message,
)
.await?
.0;
}
PendingRuntimeMessage::Insertion(insertion) => {
revision = super::engine::append_insertions(
store,
prepared,
client,
cancellation,
revision,
vec![insertion],
)
.await?
.0;
}
}
revision = append_pending(store, prepared, client, cancellation, revision, pending).await?;
}
Ok(revision)
}
@@ -174,6 +214,38 @@ enum PendingRuntimeMessage {
Insertion(MessageInsertion),
}
async fn append_pending(
store: &Store,
prepared: &PreparedRun,
client: &mut ClientPort,
cancellation: &CancellationToken,
revision: RevisionId,
pending: PendingRuntimeMessage,
) -> std::result::Result<RevisionId, RunOutcome> {
match pending {
PendingRuntimeMessage::Message(message) => Ok(super::engine::append_runtime_message(
store,
prepared,
client,
cancellation,
revision,
message,
)
.await?
.0),
PendingRuntimeMessage::Insertion(insertion) => Ok(super::engine::append_insertions(
store,
prepared,
client,
cancellation,
revision,
vec![insertion],
)
.await?
.0),
}
}
async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> {
client
.events
+40 -1
View File
@@ -47,10 +47,25 @@ pub struct TabSettings {
pub address: String,
}
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq, Serialize)]
pub struct DesktopSettings {
#[serde(default)]
pub silent_start: bool,
#[serde(default = "default_true")]
pub show_dock_icon: bool,
}
impl Default for DesktopSettings {
fn default() -> Self {
Self {
silent_start: false,
show_dock_icon: true,
}
}
}
fn default_true() -> bool {
true
}
impl TabSettings {
@@ -324,6 +339,30 @@ mod tests {
assert_eq!(store.port_settings().await.unwrap(), settings);
}
#[tokio::test]
async fn desktop_settings_show_the_dock_icon_by_default_and_round_trip() {
let store = Store::connect("sqlite::memory:").await.unwrap();
assert_eq!(
store.desktop_settings().await.unwrap(),
DesktopSettings::default()
);
assert_eq!(
serde_json::from_str::<DesktopSettings>(r#"{"silent_start":true}"#).unwrap(),
DesktopSettings {
silent_start: true,
show_dock_icon: true,
}
);
let settings = DesktopSettings {
silent_start: true,
show_dock_icon: false,
};
store.set_desktop_settings(settings).await.unwrap();
assert_eq!(store.desktop_settings().await.unwrap(), settings);
}
#[tokio::test]
async fn proxy_settings_are_write_only_and_preserve_an_unchanged_password() {
let store = Store::connect("sqlite::memory:").await.unwrap();
+657 -4
View File
@@ -5,11 +5,15 @@ mod fixtures;
use std::sync::Arc;
use bytes::Bytes;
use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{connect, proto::agent::v1 as pb},
cursor::{CursorCommand, CursorSessionRegistry},
model::{ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind},
model::{
ConversationId, ModelConfigInput, ModelSpec, ModelType, PreparedRun, PromptSpec, RunAction,
RunId, RunKind, Usage, OPENAI_CHAT_ENDPOINT,
},
provider::{FinishReason, ModelEvent},
run::RunRegistry,
store::RunStatus,
@@ -428,6 +432,472 @@ async fn injected_user_context_restarts_only_the_active_model_cycle() {
);
}
#[tokio::test]
async fn injected_user_context_aborts_pending_tools_and_ignores_late_results() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(tool_response("call-1", "Read", "{\"path\":\"/tmp/a\"}"));
let release = provider.push_gated(text_response("continued after tool interruption"));
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("interrupt-tool-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"interrupt-tool-request",
"interrupt-tool-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let exec_id = wait_for_exec(&handle, &mut output, &mut append_seqno, "Read").await;
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_injection_for(
"tool-injection",
"interrupt-tool-request",
)),
})
.await
.unwrap();
append_seqno += 1;
let mut saw_abort = false;
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
while provider.requests().len() < 2 || !saw_abort {
assert!(
tokio::time::Instant::now() < deadline,
"root model did not restart after tool interruption"
);
if let Ok(Some(frame)) =
tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await
{
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) =
server.message
{
if let Some(pb::exec_server_control_message::Message::Abort(abort)) =
control.message
{
assert_eq!(abort.id, exec_id);
saw_abort = true;
}
}
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(read_success(exec_id)),
})
.await
.unwrap();
append_seqno += 1;
release.notify_one();
drain_successfully(&handle, &mut output, &mut append_seqno).await;
let requests = provider.requests();
assert_eq!(
requests[0].history,
requests[1].history[..requests[0].history.len()]
);
let history = serde_json::to_string(&requests[1].history).unwrap();
let interrupted = history
.find("Tool execution was interrupted by a newer user message.")
.expect("interrupted tool result missing from provider history");
let injected = history
.find("injected follow-up")
.expect("injected message missing from provider history");
assert!(interrupted < injected);
}
#[tokio::test]
async fn injected_user_context_detaches_subagents_without_cancelling_them() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(tool_response(
"task-call",
"Task",
&serde_json::json!({
"description": "Inspect protocol",
"prompt": "Inspect the protocol",
"subagent_type": "generalPurpose",
"run_in_background": false
})
.to_string(),
));
let release = provider.push_gated(text_response("continued while subagent runs"));
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("detach-subagent-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"detach-subagent-request",
"detach-subagent-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let exec_id = wait_for_exec(&handle, &mut output, &mut append_seqno, "Task").await;
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_injection_for(
"subagent-injection",
"detach-subagent-request",
)),
})
.await
.unwrap();
append_seqno += 1;
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
while provider.requests().len() < 2 {
assert!(
tokio::time::Instant::now() < deadline,
"root model did not restart while subagent remained active"
);
if let Ok(Some(frame)) =
tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await
{
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) =
server.message
{
if let Some(pb::exec_server_control_message::Message::Abort(abort)) =
control.message
{
assert_ne!(abort.id, exec_id, "Task must not be aborted by injection");
}
}
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(subagent_success(exec_id)),
})
.await
.unwrap();
append_seqno += 1;
release.notify_one();
drain_successfully(&handle, &mut output, &mut append_seqno).await;
let history = serde_json::to_string(&provider.requests()[1].history).unwrap();
assert!(history.contains("Tool execution was interrupted by a newer user message."));
assert!(history.contains("injected follow-up"));
}
#[tokio::test]
async fn injected_user_context_interrupts_automatic_compaction() {
let (_directory, store) = fixtures::temp_store().await;
let model = store
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Test Model".into(),
model_id: "test-model".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: Some(10_001),
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(text_response("seed answer"));
provider.push_pending();
provider.push(text_response("continued after compacting injection"));
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let seed_state = run_to_end(
&registry,
"seed-request",
client_run_for_model(
"seed-request",
"compaction-injection-conversation",
&model.model_hash,
),
)
.await;
let handle = registry
.get_or_create("inject-during-compaction")
.await
.unwrap();
let mut output = handle.subscribe();
let mut compacting_request = client_run_for_model_with_state(
"inject-during-compaction",
"compaction-injection-conversation",
&model.model_hash,
Some(seed_state),
);
let Some(pb::agent_client_message::Message::RunRequest(request)) =
compacting_request.message.as_mut()
else {
panic!("expected RunRequest")
};
request.requested_model.as_mut().unwrap().parameters.push(
pb::requested_model::ModelParameterValue {
id: "context".into(),
value: "10001".into(),
},
);
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(compacting_request),
})
.await
.unwrap();
let mut append_seqno = 1;
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
while provider.requests().len() < 2 {
assert!(
tokio::time::Instant::now() < deadline,
"automatic compaction did not start"
);
if let Ok(Some(frame)) =
tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await
{
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_injection_for(
"compaction-injection",
"inject-during-compaction",
)),
})
.await
.unwrap();
append_seqno += 1;
let mut saw_continued = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before successful EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message {
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
saw_continued |= delta.text.contains("continued after compacting injection");
}
}
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
let requests = provider.requests();
assert_eq!(requests.len(), 3);
assert!(requests[1]
.prompt
.instructions
.starts_with("Summarize the conversation for the next model turn."));
assert!(!serde_json::to_string(&requests[1].history)
.unwrap()
.contains("injected follow-up"));
assert!(serde_json::to_string(&requests[2].history)
.unwrap()
.contains("injected follow-up"));
assert!(saw_continued);
}
#[tokio::test]
async fn stale_context_injection_is_rejected_without_failing_the_active_run() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
let release = provider.push_gated(vec![
ModelEvent::Start {
model_call_id: "active-cycle".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("active run completed".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("active-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"active-request",
"stale-injection-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
while provider.requests().is_empty() {
assert!(
tokio::time::Instant::now() < deadline,
"provider did not start"
);
if let Ok(Some(frame)) =
tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await
{
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_injection_for("stale-injection", "replaced-request")),
})
.await
.unwrap();
append_seqno += 1;
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_injection_for("stale-injection", "replaced-request")),
})
.await
.unwrap();
append_seqno += 1;
let mut rejection_count = 0;
let mut released = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before successful EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
let rejected = match server.message {
Some(pb::agent_server_message::Message::InteractionUpdate(pb::InteractionUpdate {
message:
Some(pb::interaction_update::Message::ContextInjectionState(
pb::ContextInjectionStateUpdate {
injection_id,
state:
Some(pb::ContextInjectionState {
state:
Some(pb::context_injection_state::State::Rejected(rejected)),
}),
},
)),
..
})) if injection_id == "stale-injection" => {
assert_eq!(
rejected.reason,
"InjectContextAction expected run replaced-request, active run is active-request"
);
true
}
_ => false,
};
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
if rejected {
rejection_count += 1;
if !released {
released = true;
release.notify_one();
}
}
}
assert!(released, "stale injection was not rejected");
assert_eq!(rejection_count, 1);
assert_eq!(provider.requests().len(), 1);
}
#[tokio::test]
async fn cancel_subagent_action_aborts_the_target_task_and_keeps_the_parent_running() {
let (_directory, store) = fixtures::temp_store().await;
@@ -614,6 +1084,23 @@ fn client_run() -> pb::AgentClientMessage {
}
fn client_run_for(request_id: &str, conversation_id: &str) -> pb::AgentClientMessage {
client_run_for_model(request_id, conversation_id, "test-model")
}
fn client_run_for_model(
request_id: &str,
conversation_id: &str,
model_id: &str,
) -> pb::AgentClientMessage {
client_run_for_model_with_state(request_id, conversation_id, model_id, None)
}
fn client_run_for_model_with_state(
request_id: &str,
conversation_id: &str,
model_id: &str,
state: Option<pb::ConversationStateStructure>,
) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
@@ -634,15 +1121,177 @@ fn client_run_for(request_id: &str, conversation_id: &str) -> pb::AgentClientMes
conversation_id: Some(conversation_id.into()),
run_id: Some(request_id.into()),
requested_model: Some(pb::RequestedModel {
model_id: "test-model".into(),
model_id: model_id.into(),
..Default::default()
}),
conversation_state: state,
..Default::default()
},
)),
}
}
fn text_response(text: &str) -> Vec<ModelEvent> {
vec![
ModelEvent::Start {
model_call_id: format!("call-{text}"),
},
ModelEvent::TextStart,
ModelEvent::TextDelta(text.into()),
ModelEvent::TextEnd,
ModelEvent::Usage(Usage {
input_tokens: Some(1),
output_tokens: Some(1),
total_tokens: Some(2),
..Default::default()
}),
ModelEvent::Done(FinishReason::Stop),
]
}
fn tool_response(call_id: &str, name: &str, arguments: &str) -> Vec<ModelEvent> {
vec![
ModelEvent::Start {
model_call_id: format!("call-{call_id}"),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: call_id.into(),
name: name.into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: arguments.into(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]
}
async fn wait_for_exec(
handle: &cursor_server::cursor::CursorSessionHandle,
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
append_seqno: &mut i64,
tool: &str,
) -> u32 {
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Exec");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = server.message {
let matches = match exec.message.as_ref() {
Some(pb::exec_server_message::Message::ReadArgs(_)) => tool == "Read",
Some(pb::exec_server_message::Message::SubagentArgs(_)) => tool == "Task",
_ => false,
};
if matches {
return exec.id;
}
}
acknowledge_kv(handle, append_seqno, &frame).await;
}
}
async fn drain_successfully(
handle: &cursor_server::cursor::CursorSessionHandle,
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
append_seqno: &mut i64,
) {
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before successful EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
return;
}
acknowledge_kv(handle, append_seqno, &frame).await;
}
}
fn read_success(id: u32) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
path: "/tmp/a".into(),
total_lines: 1,
file_size: 1,
output: Some(pb::read_success::Output::Content("late".into())),
..Default::default()
})),
},
)),
..Default::default()
},
)),
}
}
fn subagent_success(id: u32) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::SubagentResult(
pb::SubagentResult {
result: Some(pb::subagent_result::Result::Success(pb::SubagentSuccess {
agent_id: "detached-child".into(),
..Default::default()
})),
},
)),
..Default::default()
},
)),
}
}
async fn run_to_end(
registry: &CursorSessionRegistry,
request_id: &str,
request: pb::AgentClientMessage,
) -> pb::ConversationStateStructure {
let handle = registry.get_or_create(request_id).await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(request),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut state = None;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before EndStream");
let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
return state.expect("Run ended without a checkpoint");
}
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(update)) =
server.message
{
state = Some(update);
}
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
}
}
async fn acknowledge_kv(
handle: &cursor_server::cursor::CursorSessionHandle,
append_seqno: &mut i64,
@@ -700,13 +1349,17 @@ fn runtime_user_message() -> pb::AgentClientMessage {
}
fn runtime_injection() -> pb::AgentClientMessage {
runtime_injection_for("injection-1", "inject-request")
}
fn runtime_injection_for(injection_id: &str, expected_run_id: &str) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ConversationAction(
pb::ConversationAction {
action: Some(pb::conversation_action::Action::InjectContextAction(
pb::InjectContextAction {
injection_id: "injection-1".into(),
expected_run_id: "inject-request".into(),
injection_id: injection_id.into(),
expected_run_id: expected_run_id.into(),
payload: Some(pb::inject_context_action::Payload::UserContext(
pb::UserContextInjection {
user_message: Some(pb::UserMessage {
+46 -10
View File
@@ -46,6 +46,28 @@ fn every_tool_result_is_projected_as_string_content() {
assert_eq!(string_result.content, "plain text");
}
#[test]
fn projected_tool_result_prefixes_remain_stable() {
let first = vec![named_tool_result("Grep", &"x".repeat(64 * 1024))];
let mut second = first.clone();
second.push(fixtures::user("u2", "continue"));
let projected_first = project_messages(&first).unwrap();
let projected_second = project_messages(&second).unwrap();
assert_eq!(projected_first, projected_second[..projected_first.len()]);
}
#[test]
fn unbounded_tool_results_are_not_rewritten() {
let original = "x".repeat(64 * 1024);
let projected = project_messages(&[named_tool_result("Delete", &original)]).unwrap();
let ProjectedContent::ToolResult(result) = &projected[0].content else {
panic!("expected tool result")
};
assert_eq!(result.content, original);
}
#[test]
fn assistant_text_and_thinking_remain_separate_during_projection() {
let messages = vec![CanonicalMessage {
@@ -117,7 +139,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
.as_path(),
)
.unwrap();
assert_eq!(assets.mode(Mode::Agent).tools.len(), 22);
assert_eq!(assets.mode(Mode::Agent).tools.len(), 21);
assert_eq!(
assets
.mode(Mode::Agent)
@@ -141,7 +163,6 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"Glob",
"AskQuestion",
"Task",
"AwaitShell",
"GetMcpTools",
"FetchMcpResource",
"SwitchMode",
@@ -172,7 +193,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"SembleSearch",
"SembleFindRelated",
],
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
"98bb57a9ade7f1a572c5c5fe77a905a129d28ecfd42b8d318250f6486b09e1ec",
);
assert_mode(
&assets,
@@ -194,7 +215,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"SembleSearch",
"SembleFindRelated",
],
"e2eb8a1ebd70d53b1b2eb6bedabdce62ff070a05a6168216013d0a1144ed8bb5",
"9a7e0f9e0bd8ef0af01032fa311686f72c42ec260e3057f6fae5e68f5ed36fb8",
);
assert_mode(
&assets,
@@ -218,7 +239,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"SembleSearch",
"SembleFindRelated",
],
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
"98bb57a9ade7f1a572c5c5fe77a905a129d28ecfd42b8d318250f6486b09e1ec",
);
assert_mode(
&assets,
@@ -244,7 +265,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"SembleSearch",
"SembleFindRelated",
],
"25f7b559941baabfc9b1046455b04ca812fc41a6878ad55a43d83f0bd18cd92f",
"976b309dd91e314d4916439ebb9da8995751d011532e39934a1da7593dc78ccb",
);
assert_mode(
&assets,
@@ -263,7 +284,6 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"Write",
"Read",
"Glob",
"AwaitShell",
"GetMcpTools",
"FetchMcpResource",
"SwitchMode",
@@ -272,7 +292,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
"SembleSearch",
"SembleFindRelated",
],
"48c8e0fe825f9c2450307ca5e70cde7077c4282c135b2cd15338bd4bd0c43636",
"6de1ee86a131ca093c7143f54fffcba2fc14b32ff45fd6f5e0df1347058ad744",
);
assert_mode(
&assets,
@@ -282,7 +302,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
);
assert_eq!(
schema_digest(&assets.mode(Mode::Agent).tools),
"e53a72c1d131ff3f65c619799232440b064e90e99f5d3fcceb63e32598d3a0fc"
"282a1dff7957090d0a75eac4a46474ac7cffa1b0937bdf97354544e729bb15c2"
);
let task = assets
.mode(Mode::Agent)
@@ -519,7 +539,6 @@ fn subagent_uses_the_agent_prompt_and_only_the_captured_tool_delta() {
"Write",
"Read",
"Glob",
"AwaitShell",
"GetMcpTools",
"FetchMcpResource",
"SwitchMode",
@@ -569,6 +588,23 @@ fn tool_result_with_call(
}
}
fn named_tool_result(name: &str, output: &str) -> CanonicalMessage {
CanonicalMessage {
message_id: format!("result-{name}"),
role: Role::Tool,
origin: Origin::Tool,
content: MessageContent::ToolResult(ToolResultContent {
call_id: format!("call-{name}"),
name: name.into(),
content: output.into(),
is_error: false,
image: None,
provider_parts: Vec::new(),
}),
runtime_event_id: None,
}
}
fn assistant_tool_pair(
id: &str,
tool_round_id: &str,
-83
View File
@@ -745,89 +745,6 @@ async fn unknown_exec_id_is_a_protocol_error() {
));
}
#[tokio::test]
async fn await_shell_consumes_the_background_output_file_terminal_state() {
let runtime = CursorToolRuntime::default();
let dispatcher = ToolDispatcher::new(runtime.clone());
let mut await_call = call("await-call", "AwaitShell");
await_call.arguments = json!({
"shell_id": "42",
"block_until_ms": 1000,
"pattern": "ready",
});
await_call.arguments_text = await_call.arguments.to_string();
let completed = HashSet::new();
let started = HashSet::new();
let dispatched = dispatcher
.start_batch(
&[await_call],
ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
},
&[],
&BTreeMap::new(),
&exec_context(),
)
.await
.unwrap();
let exec = dispatched[0]
.messages
.iter()
.find_map(|message| match message.message.as_ref() {
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => Some(exec),
_ => None,
})
.unwrap();
let Some(pb::exec_server_message::Message::ReadArgs(read)) = exec.message.as_ref() else {
panic!("expected AwaitShell ReadArgs")
};
assert_eq!(read.path, "/tmp/terminals/42.txt");
let event = codec::client_event(
&pb::ExecClientMessage {
id: exec.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(
"server ready\nexit_code: 0\n".into(),
)),
..Default::default()
})),
},
)),
..Default::default()
},
&runtime,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = event else {
panic!("expected completed AwaitShell")
};
assert_eq!(completion.result().call_id, "await-call");
assert!(!completion.result().is_error);
let Some(pb::tool_call::Tool::AwaitToolCall(tool)) = completion.tool_call().tool.as_ref()
else {
panic!("expected AwaitToolCall")
};
let pb::await_result::Result::Success(success) =
tool.result.as_ref().unwrap().result.as_ref().unwrap()
else {
panic!("expected Await success")
};
let pb::await_success::AwaitResult::Complete(complete) = success.await_result.as_ref().unwrap()
else {
panic!("expected completed background task")
};
assert_eq!(complete.task_id, "42");
assert_eq!(complete.exit_code, Some(0));
assert_eq!(complete.regex_match.as_deref(), Some("ready"));
}
#[tokio::test]
async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
let (directory, store) = fixtures::temp_store().await;