fix: repair stream lifecycle reliability

This commit is contained in:
leookun
2026-08-28 13:00:39 +08:00
parent 370f120b8a
commit 59379f1f72
7 changed files with 108 additions and 106 deletions
+1 -1
View File
@@ -21,7 +21,7 @@ use crate::{
store::Store, store::Store,
}; };
use super::{inbox::OrderedInbox, CursorCommand, CursorSessionHandle}; use super::{inbox::OrderedInbox, lifecycle, CursorCommand, CursorSessionHandle};
pub struct CursorActor; pub struct CursorActor;
+1
View File
@@ -42,6 +42,7 @@ pub fn fail(handle: &CursorSessionHandle, error: &Error) -> Result<()> {
} }
pub fn cancel(handle: &CursorSessionHandle) -> Result<()> { pub fn cancel(handle: &CursorSessionHandle) -> Result<()> {
handle.cancel();
// Always close the output even if encoding fails, to prevent silent hangs. // Always close the output even if encoding fails, to prevent silent hangs.
match encode_error_end_stream(&ConnectStreamError { match encode_error_end_stream(&ConnectStreamError {
code: ConnectCode::Canceled, code: ConnectCode::Canceled,
+2
View File
@@ -316,6 +316,8 @@ impl CursorSessionRegistry {
// Create the notification future BEFORE checking state to avoid // Create the notification future BEFORE checking state to avoid
// a race where a notification fires between state check and await. // a race where a notification fires between state check and await.
let changed = self.inner.route_changed.notified(); let changed = self.inner.route_changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
if self.inner.runs.lock().await.contains_key(request_id) { if self.inner.runs.lock().await.contains_key(request_id) {
return CursorRoute::Local; return CursorRoute::Local;
} }
+21 -3
View File
@@ -54,8 +54,10 @@ pub(super) async fn start(
let call = normalized_call.as_ref().unwrap_or(call); let call = normalized_call.as_ref().unwrap_or(call);
match normalized(&call.name).as_str() { match normalized(&call.name).as_str() {
"shell" | "bash" | "read" | "delete" | "grep" | "glob" | "readlints" | "task" | "callmcptool" "shell" | "bash" | "read" | "delete" | "grep" | "glob" | "readlints" | "task"
| "fetchmcpresource" | "getmcptools" => exec::start(runtime, call, context).await, | "callmcptool" | "fetchmcpresource" | "getmcptools" => {
exec::start(runtime, call, context).await
}
"write" | "strreplace" | "editnotebook" => edit::start(runtime, call, context).await, "write" | "strreplace" | "editnotebook" => edit::start(runtime, call, context).await,
"askquestion" | "websearch" | "webfetch" | "switchmode" | "createplan" "askquestion" | "websearch" | "webfetch" | "switchmode" | "createplan"
| "generateimage" => interaction::start(runtime, call).await, | "generateimage" => interaction::start(runtime, call).await,
@@ -66,7 +68,7 @@ pub(super) async fn start(
} }
fn normalize_block_until_ms(call: &ToolCall) -> Result<Option<ToolCall>> { fn normalize_block_until_ms(call: &ToolCall) -> Result<Option<ToolCall>> {
if normalized(&call.name) != "shell" { if !is_shell_tool(&call.name) {
return Ok(None); return Ok(None);
} }
let Some(value) = call.arguments.get("block_until_ms") else { let Some(value) = call.arguments.get("block_until_ms") else {
@@ -133,6 +135,10 @@ pub(super) async fn resume_interaction(
interaction::resume(results, search, fetch, pending, response).await interaction::resume(results, search, fetch, pending, response).await
} }
fn is_shell_tool(name: &str) -> bool {
matches!(normalized(name).as_str(), "shell" | "bash")
}
pub(super) fn normalized(name: &str) -> String { pub(super) fn normalized(name: &str) -> String {
name.chars() name.chars()
.filter(|character| character.is_ascii_alphanumeric()) .filter(|character| character.is_ascii_alphanumeric())
@@ -167,6 +173,18 @@ mod tests {
assert_eq!(call.arguments["block_until_ms"].as_i64(), Some(45_000)); assert_eq!(call.arguments["block_until_ms"].as_i64(), Some(45_000));
} }
#[test]
fn bash_accepts_integer_valued_float_timeout() {
let call = tool(
"Bash",
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] #[test]
fn shell_rejects_fractional_timeout() { fn shell_rejects_fractional_timeout() {
let call = tool( let call = tool(
+21 -51
View File
@@ -15,29 +15,6 @@ use crate::{
Error, Result, Error, Result,
}; };
macro_rules! dbg_log {
($loc:expr, $msg:expr, $data:expr) => {
{
let _ = (|| -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>> {
use std::io::Write;
let payload = serde_json::json!({
"sessionId": "216d24",
"location": $loc,
"message": $msg,
"data": $data,
"timestamp": std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH)?.as_millis() as u64,
});
let mut f = std::fs::OpenOptions::new().create(true).append(true).open(
concat!(env!("CARGO_MANIFEST_DIR"), "/../.cursor/debug-216d24.log")
)?;
writeln!(f, "{}", payload)?;
f.flush()?;
Ok(())
})();
}
};
}
use super::{ use super::{
apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, CallRecorder, apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, CallRecorder,
FinishReason, ModelEvent, Provider, ProviderStream, FinishReason, ModelEvent, Provider, ProviderStream,
@@ -84,13 +61,13 @@ impl Provider for OpenAiChatProvider {
let recorder = self.recorder.clone(); let recorder = self.recorder.clone();
Box::pin(try_stream! { Box::pin(try_stream! {
let ModelInvocation { call_id, request, .. } = invocation; let ModelInvocation { call_id, request, .. } = invocation;
dbg_log!("openai_chat.rs:stream", "Provider stream started", serde_json::json!({ tracing::debug!(
"model": request.model.model_id, model = %request.model.model_id,
"url": config.request_url, call_id = %call_id,
"call_id": call_id.clone(), history_len = request.history.len(),
"history_len": request.history.len(), tools_count = request.prompt.tools.len(),
"tools_count": request.prompt.tools.len() "OpenAI Chat provider stream started"
})); );
let messages = openai_chat_messages(&request.prompt.instructions, &request.history)?; let messages = openai_chat_messages(&request.prompt.instructions, &request.history)?;
let mut body = json!({ let mut body = json!({
"model": request.model.model_id, "model": request.model.model_id,
@@ -117,17 +94,11 @@ impl Provider for OpenAiChatProvider {
}; };
let response = match response { let response = match response {
Ok(r) => { Ok(r) => {
dbg_log!("openai_chat.rs:stream", "HTTP response received", serde_json::json!({ tracing::debug!(status = r.status().as_u16(), "OpenAI Chat HTTP response received");
"status": r.status().as_u16(),
"url": config.request_url
}));
r r
} }
Err(e) => { Err(e) => {
dbg_log!("openai_chat.rs:stream", "HTTP request FAILED", serde_json::json!({ tracing::debug!(error = %e, "OpenAI Chat HTTP request failed");
"error": e.to_string(),
"url": config.request_url
}));
Err(Error::from(e))? Err(Error::from(e))?
} }
}; };
@@ -168,28 +139,27 @@ impl Provider for OpenAiChatProvider {
loop_iteration += 1; loop_iteration += 1;
let event = tokio::select! { let event = tokio::select! {
_ = cancellation.cancelled() => { _ = cancellation.cancelled() => {
dbg_log!("openai_chat.rs:stream", "Stream cancelled by token", serde_json::json!({ tracing::debug!(
"iteration": loop_iteration, iteration = loop_iteration,
"saw_done_marker": saw_done_marker, saw_done_marker,
"tool_count": tools.len() tool_count = tools.len(),
})); "OpenAI Chat stream cancelled"
);
return; return;
} }
event = source.next() => event, event = source.next() => event,
}; };
let Some(event) = event else { let Some(event) = event else {
dbg_log!("openai_chat.rs:stream", "SSE stream ended (None)", serde_json::json!({ tracing::debug!(
"iteration": loop_iteration, iteration = loop_iteration,
"saw_done_marker": saw_done_marker saw_done_marker,
})); "OpenAI Chat SSE stream ended"
);
break; break;
}; };
let event = event.map_err(|error| { let event = event.map_err(|error| {
let err_msg = error.to_string(); let err_msg = error.to_string();
dbg_log!("openai_chat.rs:stream", "SSE event error", serde_json::json!({ tracing::debug!(iteration = loop_iteration, error = %error, "OpenAI Chat SSE event failed");
"iteration": loop_iteration,
"error": err_msg.clone()
}));
Error::Provider(format!("OpenAI Chat SSE: {err_msg}")) Error::Provider(format!("OpenAI Chat SSE: {err_msg}"))
})?; })?;
if event.data == "[DONE]" { saw_done_marker = true; break; } if event.data == "[DONE]" { saw_done_marker = true; break; }
+30 -51
View File
@@ -11,29 +11,6 @@ use crate::{
Error, Result, Error, Result,
}; };
macro_rules! dbg_log {
($loc:expr, $msg:expr, $data:expr) => {
{
let _ = (|| -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>> {
use std::io::Write;
let payload = serde_json::json!({
"sessionId": "216d24",
"location": $loc,
"message": $msg,
"data": $data,
"timestamp": std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH)?.as_millis() as u64,
});
let mut f = std::fs::OpenOptions::new().create(true).append(true).open(
concat!(env!("CARGO_MANIFEST_DIR"), "/../.cursor/debug-216d24.log")
)?;
writeln!(f, "{}", payload)?;
f.flush()?;
Ok(())
})();
}
};
}
use super::{ use super::{
normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider, normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider,
OpenAiResponsesProvider, Provider, ProviderStream, OpenAiResponsesProvider, Provider, ProviderStream,
@@ -114,12 +91,12 @@ impl Provider for ProviderRouter {
let stream_cancellation = cancellation.clone(); let stream_cancellation = cancellation.clone();
let mut stream = provider.stream(invocation, cancellation); let mut stream = provider.stream(invocation, cancellation);
let stream_started = std::time::Instant::now(); let stream_started = std::time::Instant::now();
dbg_log!("router.rs:stream", "provider stream created", serde_json::json!({ tracing::debug!(
"model": selected, model = %selected,
"provider_type": format!("{:?}", provider_type), provider_type = ?provider_type,
"url": config.request_url, timeout_ms = config.request_timeout.as_millis() as u64,
"timeout_ms": config.request_timeout.as_millis() as u64, "provider stream created"
})); );
let mut last_event_time = std::time::Instant::now(); let mut last_event_time = std::time::Instant::now();
let mut event_count: u64 = 0; let mut event_count: u64 = 0;
while let Some(event) = stream.next().await { while let Some(event) = stream.next().await {
@@ -132,37 +109,39 @@ impl Provider for ProviderRouter {
let event_name = match &event { let event_name = match &event {
super::ModelEvent::Start { .. } => "Start", super::ModelEvent::Start { .. } => "Start",
super::ModelEvent::TextStart => "TextStart", super::ModelEvent::TextStart => "TextStart",
super::ModelEvent::TextDelta(d) => { dbg_log!("router.rs:stream", "TextDelta", serde_json::json!({"gap_ms": gap_ms, "elapsed_ms": elapsed_ms, "delta_len": d.len(), "event_count": event_count})); "TextDelta" }, super::ModelEvent::TextDelta(_) => "TextDelta",
super::ModelEvent::TextEnd => "TextEnd", super::ModelEvent::TextEnd => "TextEnd",
super::ModelEvent::ThinkingStart => "ThinkingStart", super::ModelEvent::ThinkingStart => "ThinkingStart",
super::ModelEvent::ThinkingDelta(d) => { dbg_log!("router.rs:stream", "ThinkingDelta", serde_json::json!({"gap_ms": gap_ms, "elapsed_ms": elapsed_ms, "delta_len": d.len(), "event_count": event_count})); "ThinkingDelta" }, super::ModelEvent::ThinkingDelta(_) => "ThinkingDelta",
super::ModelEvent::ThinkingEnd => "ThinkingEnd", super::ModelEvent::ThinkingEnd => "ThinkingEnd",
super::ModelEvent::ToolCallStart { name, .. } => { dbg_log!("router.rs:stream", "ToolCallStart", serde_json::json!({"name": name, "elapsed_ms": elapsed_ms})); "ToolCallStart" }, super::ModelEvent::ToolCallStart { .. } => "ToolCallStart",
super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta", super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta",
super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd", super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd",
super::ModelEvent::ProviderReplayState(_) => "ReplayState", super::ModelEvent::ProviderReplayState(_) => "ReplayState",
super::ModelEvent::Usage(_) => "Usage", super::ModelEvent::Usage(_) => "Usage",
super::ModelEvent::Done(reason) => { dbg_log!("router.rs:stream", "Done", serde_json::json!({"reason": format!("{:?}", reason), "elapsed_ms": elapsed_ms, "event_count": event_count})); "Done" }, super::ModelEvent::Done(_) => "Done",
}; };
if gap_ms > 5000 { if gap_ms > 5000 {
dbg_log!("router.rs:stream", "SLOW GAP detected between events", serde_json::json!({ tracing::debug!(
"gap_ms": gap_ms, gap_ms,
"elapsed_ms": elapsed_ms, elapsed_ms,
"event": event_name, event = event_name,
"event_count": event_count, event_count,
})); "slow gap detected between provider events"
);
} }
recorder.event(&event).await?; recorder.event(&event).await?;
last_event_time = now; last_event_time = now;
yield event; yield event;
} }
Err(error) => { Err(error) => {
dbg_log!("router.rs:stream", "PROVIDER ERROR", serde_json::json!({ tracing::debug!(
"error": error.to_string(), error = %error,
"elapsed_ms": elapsed_ms, elapsed_ms,
"gap_ms": gap_ms, gap_ms,
"event_count": event_count, event_count,
})); "provider stream error"
);
recorder.failed(&error).await?; recorder.failed(&error).await?;
Err(error)?; Err(error)?;
} }
@@ -171,15 +150,15 @@ impl Provider for ProviderRouter {
if !recorder.is_finished() { if !recorder.is_finished() {
let elapsed_ms = stream_started.elapsed().as_millis() as u64; let elapsed_ms = stream_started.elapsed().as_millis() as u64;
if stream_cancellation.is_cancelled() { if stream_cancellation.is_cancelled() {
dbg_log!("router.rs:stream", "stream ended: cancelled", serde_json::json!({"elapsed_ms": elapsed_ms, "event_count": event_count})); tracing::debug!(elapsed_ms, event_count, "provider stream ended after cancellation");
recorder.cancelled().await?; recorder.cancelled().await?;
} else { } else {
let error = Error::Provider("provider stream ended without Done".into()); let error = Error::Provider("provider stream ended without Done".into());
dbg_log!("router.rs:stream", "stream ended: WITHOUT DONE (fatal)", serde_json::json!({ tracing::warn!(
"elapsed_ms": elapsed_ms, elapsed_ms,
"event_count": event_count, event_count,
"error": "provider stream ended without Done", "provider stream ended without Done"
})); );
recorder.failed(&error).await?; recorder.failed(&error).await?;
Err(error)?; Err(error)?;
} }
+32
View File
@@ -19,6 +19,38 @@ use cursor_server::{
}; };
use prost::Message; use prost::Message;
#[tokio::test]
async fn abort_command_cancels_the_run_and_closes_output() {
let (_directory, store) = fixtures::temp_store().await;
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(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("abort-request").await.unwrap();
let mut output = handle.subscribe();
handle.command(CursorCommand::Abort).await.unwrap();
let frame = tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
.await
.unwrap()
.expect("Abort must emit a terminal frame");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags, connect::END_STREAM_FLAG);
let payload: serde_json::Value = serde_json::from_slice(&payload).unwrap();
assert_eq!(payload["error"]["code"], "canceled");
assert!(handle.cancellation().is_cancelled());
assert_eq!(output.recv().await, None);
}
#[tokio::test] #[tokio::test]
async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_error() { async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_error() {
let (_directory, store) = fixtures::temp_store().await; let (_directory, store) = fixtures::temp_store().await;