diff --git a/server/src/api/cursor/bidi.rs b/server/src/api/cursor/bidi.rs index 8592c66..f069981 100644 --- a/server/src/api/cursor/bidi.rs +++ b/server/src/api/cursor/bidi.rs @@ -184,7 +184,11 @@ pub async fn append( request: DecodedAppend, parent: Option, ) -> Result { - let handle = registry.get_or_create(&request.request_id).await?; + let replace_closing = request.model_id().is_some(); + let handle = registry + .get_or_create_for_append(&request.request_id, replace_closing) + .await?; + let _admission = handle.admit()?; if let Some(conversation_id) = request.conversation_id() { handle.set_conversation_id(conversation_id)?; } diff --git a/server/src/api/cursor/handlers.rs b/server/src/api/cursor/handlers.rs index aa0c0d3..5d78278 100644 --- a/server/src/api/cursor/handlers.rs +++ b/server/src/api/cursor/handlers.rs @@ -19,9 +19,7 @@ use crate::{ connect, proto::{agent::v1 as agent, aiserver::v1 as ai}, }, - services::{ - account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab, - }, + services::{account, analytics, knowledge, model_catalog, tab}, transport::{TransportParent, TransportRegistry}, }, Result, @@ -123,16 +121,13 @@ async fn run_sse_handler( let (parts, body) = buffered(request).await?; let request: agent::BidiRequestId = connect::decode_unary(&body)?; let route = registry.wait_route(&request.request_id).await; - let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await; - if let Some(trace) = &trace { - trace - .request( - "run_sse_request", - &body, - serde_json::json!({"request_id": request.request_id}), - ) - .await; - } + let trace = registry.trace(&request.request_id); + trace.resume(); + trace.request( + "run_sse_request", + body.clone(), + serde_json::json!({"request_id": request.request_id}), + ); match route { crate::cursor::transport::TransportRoute::Local => { run_sse::stream(®istry, &request.request_id).await @@ -143,7 +138,14 @@ async fn run_sse_handler( Request::from_parts(parts, Body::from(body)), ) .await?; - Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await) + Ok(run_sse::upstream( + registry, + request.request_id, + generation, + response, + Some(trace), + ) + .await) } } } @@ -159,6 +161,7 @@ async fn bidi_handler( let first_model = decoded.model_id().map(str::to_owned); let conversation_id = decoded.conversation_id().map(str::to_owned); let trace_metadata = decoded.trace_metadata(); + let trace = registry.trace(&decoded.request_id); let local = if let Some(model_id) = decoded.model_id() { // 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。 if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX) @@ -183,14 +186,18 @@ async fn bidi_handler( } else if registry.upstream(&decoded.request_id).await { false } else { + trace.resume(); + trace.request( + "bidi_request", + body.clone(), + trace_outcome(trace_metadata, false, "missing_transport", None), + ); return Err(crate::Error::Protocol( "first BidiAppend message must select a model".into(), )); }; - let trace = if first_model.is_some() { - CursorTraceRecorder::begin( - registry.store().clone(), - &decoded.request_id, + if first_model.is_some() { + trace.begin( conversation_id.as_deref(), if local { "local_byok" @@ -198,26 +205,61 @@ async fn bidi_handler( "cursor_official" }, first_model.as_deref(), - ) - .await + ); } else { - CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await - }; - if let Some(trace) = &trace { - trace.request("bidi_request", &body, trace_metadata).await; + trace.resume(); } if !local { if first_model.is_some() { registry.mark_upstream(&decoded.request_id).await; } + trace.request( + "bidi_request", + body.clone(), + trace_outcome(trace_metadata, true, "upstream", None), + ); return proxy::forward( Extension(proxy), Request::from_parts(parts, Body::from(body)), ) .await; } - let parent = parent_headers(&parts.headers)?; - bidi::append(®istry, decoded, parent).await?; + let parent = match parent_headers(&parts.headers) { + Ok(parent) => parent, + Err(error) => { + trace.request( + "bidi_request", + body, + trace_outcome( + trace_metadata, + false, + "invalid_parent", + Some(error.to_string()), + ), + ); + return Err(error); + } + }; + match bidi::append(®istry, decoded, parent).await { + Ok(_) => trace.request( + "bidi_request", + body, + trace_outcome(trace_metadata, true, "local", None), + ), + Err(error) => { + trace.request( + "bidi_request", + body, + trace_outcome( + trace_metadata, + false, + "command_rejected", + Some(error.to_string()), + ), + ); + return Err(error); + } + } let mut response = Response::new(axum::body::Body::empty()); *response.status_mut() = StatusCode::OK; response.headers_mut().insert( @@ -227,6 +269,22 @@ async fn bidi_handler( Ok(response) } +fn trace_outcome( + mut metadata: serde_json::Value, + accepted: bool, + route_outcome: &str, + error: Option, +) -> serde_json::Value { + if let Some(metadata) = metadata.as_object_mut() { + metadata.insert("accepted".into(), accepted.into()); + metadata.insert("route_outcome".into(), route_outcome.into()); + if let Some(error) = error { + metadata.insert("error".into(), error.into()); + } + } + metadata +} + async fn buffered(request: Request) -> Result<(axum::http::request::Parts, Bytes)> { let (parts, body) = request.into_parts(); let body = to_bytes(body, usize::MAX) diff --git a/server/src/api/cursor/run_sse.rs b/server/src/api/cursor/run_sse.rs index b81b97f..6da2b52 100644 --- a/server/src/api/cursor/run_sse.rs +++ b/server/src/api/cursor/run_sse.rs @@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result Response { let (parts, body) = response.into_parts(); if let Some(trace) = &trace { - trace.response_started(parts.status.as_u16()).await; + trace.response_started(parts.status.as_u16()); } let stream = async_stream::stream! { let _guard = UpstreamRunGuard { @@ -180,15 +180,15 @@ impl TraceStreamSink { while let Some(event) = receiver.recv().await { match event { TraceStreamEvent::Chunk(chunk) => { - trace.response_chunk(source, &chunk).await; + trace.response_chunk(source, chunk); } TraceStreamEvent::Finish(error) => { - trace.finish(error.as_deref()).await; + trace.finish(error.as_deref()); return; } } } - trace.finish(None).await; + trace.finish(None); }); Self { sender: Some(sender), diff --git a/server/src/app.rs b/server/src/app.rs index 62cb015..6f97c86 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -66,6 +66,7 @@ impl App { plugin_runtime, plugins, clients.clone(), + config.app_version.clone(), )?; let harness = control.cursor_harness().clone(); let mut router = api::router(registry.clone(), clients)?; diff --git a/server/src/control/ads.rs b/server/src/control/ads.rs index 7749f65..bd5ffe9 100644 --- a/server/src/control/ads.rs +++ b/server/src/control/ads.rs @@ -16,7 +16,8 @@ use super::ControlService; // 此广告拉取不涉及用户隐私,用户id随机产生 // 开源项目广告为作者唯一收入来源,PR时请勿去除此广告 -pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu"; +// pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu"; +pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu"; pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID"; pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS"; pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version"; diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 28e9e09..7a9d827 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -42,6 +42,7 @@ pub struct ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, model_tests: Arc>>, } @@ -153,6 +154,7 @@ impl ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, ) -> Result { Ok(Self { cursor_harness: CursorHarness::new(store.clone())?, @@ -161,6 +163,7 @@ impl ControlService { plugin_runtime, plugins, clients, + app_version, model_tests: Arc::new(Mutex::new(BTreeMap::new())), }) } @@ -265,7 +268,7 @@ impl ControlService { .get(ADS_ENDPOINT) .header(DEVICE_ID_HEADER, installation_id) .header(OS_HEADER, std::env::consts::OS) - .header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION")) + .header(APP_VERSION_HEADER, &self.app_version) .header(LANGUAGE_HEADER, language) .timeout(std::time::Duration::from_secs(60)); if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) { @@ -299,7 +302,7 @@ impl ControlService { .post(endpoint) .header(DEVICE_ID_HEADER, installation_id) .header(OS_HEADER, std::env::consts::OS) - .header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION")) + .header(APP_VERSION_HEADER, &self.app_version) .json(input) .timeout(std::time::Duration::from_secs(5)) .send() diff --git a/server/src/cursor/checkpoint/builder.rs b/server/src/cursor/checkpoint/builder.rs index 28c1c29..704833a 100644 --- a/server/src/cursor/checkpoint/builder.rs +++ b/server/src/cursor/checkpoint/builder.rs @@ -278,19 +278,17 @@ impl CheckpointBuilder { ), }); if let Some(trace) = handle.trace() { - trace - .artifact( - "checkpoint", - "byok_server", - &checkpoint.encode_to_vec(), - serde_json::json!({ - "root_message_count": checkpoint.root_prompt_messages_json.len(), - "turn_count": checkpoint.turns.len(), - "pending_tool_call_count": checkpoint.pending_tool_calls.len(), - "emit_status": if result.is_ok() { "sent" } else { "error" }, - }), - ) - .await; + trace.artifact( + "checkpoint", + "byok_server", + &checkpoint.encode_to_vec(), + serde_json::json!({ + "root_message_count": checkpoint.root_prompt_messages_json.len(), + "turn_count": checkpoint.turns.len(), + "pending_tool_call_count": checkpoint.pending_tool_calls.len(), + "emit_status": if result.is_ok() { "sent" } else { "error" }, + }), + ); } result } diff --git a/server/src/cursor/compile/context.rs b/server/src/cursor/compile/context.rs index e29e634..859ecc5 100644 --- a/server/src/cursor/compile/context.rs +++ b/server/src/cursor/compile/context.rs @@ -12,7 +12,7 @@ use crate::{ protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer, tools::runtime::McpRoute, }, - model::ToolDefinition, + model::{normalize_tool_name, ToolDefinition}, store::BlobId, Error, Result, }; @@ -493,7 +493,7 @@ pub fn dynamic_mcp( })?), }; let parameters = normalize_mcp_parameters(&wire.name, parameters)?; - let name = model_tool_name(&wire.name); + let name = normalize_tool_name(&wire.name); let definition = ToolDefinition { name: name.clone(), description: wire.description.clone(), @@ -553,18 +553,6 @@ fn invalid_mcp_parameters(tool_name: &str) -> Error { )) } -fn model_tool_name(name: &str) -> String { - name.chars() - .map(|character| { - if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') { - character - } else { - '_' - } - }) - .collect() -} - fn prost_value(value: &prost_types::Value) -> Value { use prost_types::value::Kind; match value.kind.as_ref() { diff --git a/server/src/cursor/compile/run.rs b/server/src/cursor/compile/run.rs index f351ead..b3712ce 100644 --- a/server/src/cursor/compile/run.rs +++ b/server/src/cursor/compile/run.rs @@ -117,9 +117,7 @@ pub(crate) async fn prepare( "selected_source": "root_prompt_messages_json", }); let encoded = serde_json::to_vec(&summary)?; - trace - .artifact("history_projection", "byok_server", &encoded, summary) - .await; + trace.artifact("history_projection", "byok_server", &encoded, summary); } let mut request_context = context::hydrate(request, context_sync).await?; if let Some(rules_dir) = local_rules_dir { diff --git a/server/src/cursor/conversation/command.rs b/server/src/cursor/conversation/command.rs index 3cdd2f7..4ba5f2b 100644 --- a/server/src/cursor/conversation/command.rs +++ b/server/src/cursor/conversation/command.rs @@ -1,6 +1,19 @@ //! Defines commands accepted by a Conversation runtime. -use crate::cursor::protocol::proto::agent::v1 as pb; +use crate::{cursor::protocol::proto::agent::v1 as pb, Error}; + +#[derive(Debug)] +pub enum RunFinish { + TurnCompleted, + Transport(TransportFinish), +} + +#[derive(Debug)] +pub enum TransportFinish { + Success, + Failed(Error), + Cancelled, +} #[derive(Debug)] pub enum TransportCommand { @@ -8,6 +21,9 @@ pub enum TransportCommand { seqno: i64, message: Box, }, + RunFinished { + generation: u64, + finish: RunFinish, + }, Disconnect, - Close, } diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index 654937b..349fccc 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -34,7 +34,7 @@ use crate::{ Error, Result, }; -use super::{CompiledMessages, ConversationRegistry, MessageDelivery}; +use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish}; use crate::cursor::transport::TransportHandle; pub struct ConversationOutput { @@ -110,7 +110,7 @@ impl ConversationOutput { } } - pub async fn run(mut self) -> Result<()> { + pub async fn run(mut self) -> Result { let result = self.run_inner().await; if let Err(error) = &result { if !self.superseded.is_cancelled() { @@ -143,7 +143,7 @@ impl ConversationOutput { result } - async fn run_inner(&mut self) -> Result<()> { + async fn run_inner(&mut self) -> Result { if self.context.compacting { self.handle.emit(&events::summary_started())?; } @@ -175,7 +175,7 @@ impl ConversationOutput { if self.superseded.is_cancelled() { worker.abort(); self.abort_execs().await; - return Ok(()); + return Ok(RunFinish::Transport(TransportFinish::Cancelled)); } let input = if let Ok(action) = self.runtime_actions.try_recv() { Input::RuntimeAction(Some(Box::new(action))) @@ -187,7 +187,7 @@ impl ConversationOutput { _ = self.superseded.cancelled() => { worker.abort(); self.abort_execs().await; - return Ok(()); + return Ok(RunFinish::Transport(TransportFinish::Cancelled)); } action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)), event = self.core.events.recv() => Input::Event(event), @@ -719,7 +719,7 @@ impl ConversationOutput { if self.superseded.is_cancelled() { worker.abort(); self.abort_execs().await; - return Ok(()); + return Ok(RunFinish::Transport(TransportFinish::Cancelled)); } return match outcome { RunOutcome::Completed => { @@ -735,8 +735,7 @@ impl ConversationOutput { for _ in 0..3 { self.checkpoint.publish(&self.handle, &checkpoint).await?; } - finish_success(&self.handle); - return Ok(()); + return Ok(RunFinish::TurnCompleted); } let checkpoints = final_checkpoint.take().ok_or_else(|| { Error::Protocol("Completed without final state".into()) @@ -752,18 +751,19 @@ impl ConversationOutput { ttft_breakdown: None, message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)), })?; - finish_success(&self.handle); - Ok(()) + Ok(RunFinish::TurnCompleted) } RunOutcome::Cancelled => { worker.abort(); self.abort_execs().await; - finish_cancelled(&self.handle) + Ok(RunFinish::Transport(TransportFinish::Cancelled)) } RunOutcome::Failed(failure) => { worker.abort(); self.abort_execs().await; - finish_failed(&self.handle, &cursor_error(failure)) + Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error( + failure, + )))) } }; } diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index dc3f1ef..3a1edec 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -18,18 +18,22 @@ use crate::{ }, transport::{OrderedInbox, TransportHandle}, }, - run::{CommandResult, RunEngine, RunHandle}, + run::{CommandResult, RunEngine, RunHandle, RunPhase}, }; use super::{ CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies, - ConversationRegistry, MessageDelivery, TransportCommand, + ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish, }; pub struct ConversationRuntime; +const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2); + #[derive(Clone)] struct RunGeneration { + id: u64, + request: pb::AgentRunRequest, superseded: CancellationToken, finished: CancellationToken, run: Arc>>, @@ -41,6 +45,14 @@ struct RunGeneration { struct FinishGeneration(CancellationToken); +struct TransportActorGuard(TransportHandle); + +impl Drop for TransportActorGuard { + fn drop(&mut self) { + self.0.close_transport(); + } +} + impl Drop for FinishGeneration { fn drop(&mut self) { self.0.cancel(); @@ -54,6 +66,7 @@ impl ConversationRuntime { mut receiver: mpsc::Receiver, ) { tokio::spawn(async move { + let _actor_guard = TransportActorGuard(handle.clone()); let dependencies = registry.dependencies().clone(); let blob_sync = BlobSynchronizer::new( handle.request_id().into(), @@ -65,19 +78,70 @@ impl ConversationRuntime { let context_sync = RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone()); let mut current = None::; + let mut next_generation = 1_u64; + let mut pending_finish = None::<(u64, TransportFinish)>; + let mut draining = false; + let mut waiting_for_action = false; loop { - let command = match receiver.recv().await { - Some(command) => command, - None => { - handle.mark_disconnected(); - if let Some(generation) = current.as_ref() { - generation.superseded.cancel(); - if let Some(run) = generation.run.lock().clone() { - run.cancel(); + let command = if draining { + if !handle.admissions_drained() { + tokio::select! { + command = receiver.recv() => match command { + Some(command) => command, + None => { + finish_pending(&handle, ¤t, pending_finish.take()); + break; + } + }, + _ = handle.wait_admissions_drained() => continue, + } + } else { + handle.mark_draining(); + match receiver.try_recv() { + Ok(command) => command, + Err(mpsc::error::TryRecvError::Empty) + | Err(mpsc::error::TryRecvError::Disconnected) => { + finish_pending(&handle, ¤t, pending_finish.take()); + break; } } - super::finish_cancelled(&handle).ok(); - break; + } + } else if waiting_for_action { + tokio::select! { + command = receiver.recv() => match command { + Some(command) => command, + None => { + handle.mark_disconnected(); + super::finish_success(&handle); + break; + } + }, + _ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => { + let Some(generation) = current.as_ref() else { + super::finish_success(&handle); + break; + }; + handle.begin_close(); + pending_finish = Some((generation.id, TransportFinish::Success)); + draining = true; + waiting_for_action = false; + continue; + } + } + } else { + match receiver.recv().await { + Some(command) => command, + None => { + handle.mark_disconnected(); + if let Some(generation) = current.as_ref() { + generation.superseded.cancel(); + if let Some(run) = generation.run.lock().clone() { + run.cancel(); + } + } + super::finish_cancelled(&handle).ok(); + break; + } } }; match command { @@ -92,11 +156,41 @@ impl ConversationRuntime { let _ = handle.emit(&codec::abort(id)); } } - super::finish_cancelled(&handle).ok(); + let turn_completed = waiting_for_action + || current.as_ref().is_some_and(|generation| { + generation + .run + .lock() + .as_ref() + .is_none_or(|run| run.phase() != RunPhase::Running) + }); + if turn_completed { + super::finish_success(&handle); + } else { + super::finish_cancelled(&handle).ok(); + } break; } - TransportCommand::Close => { - break; + TransportCommand::RunFinished { generation, finish } => { + if !current + .as_ref() + .is_some_and(|current| current.id == generation) + { + continue; + } + match finish { + RunFinish::TurnCompleted => { + pending_finish = None; + draining = false; + waiting_for_action = true; + } + RunFinish::Transport(finish) => { + waiting_for_action = false; + handle.begin_close(); + pending_finish = Some((generation, finish)); + draining = true; + } + } } TransportCommand::Append { seqno, message } => { for (_seqno, message) in inbox.push(seqno, *message) { @@ -105,6 +199,12 @@ impl ConversationRuntime { Some(pb::agent_client_message::Message::RunRequest( request, )) => { + waiting_for_action = false; + if draining { + handle.reopen(); + draining = false; + pending_finish = None; + } if let Some(conversation_id) = request.conversation_id.as_deref() { @@ -117,60 +217,21 @@ impl ConversationRuntime { "invalid Cursor conversation id" ); let _ = super::finish_failed(&handle, &error); - let _ = - handle.command(TransportCommand::Close).await; return; } } - let previous_finished = - if let Some(previous) = current.take() { - previous.superseded.cancel(); - if let Some(run) = previous.run.lock().clone() { - run.cancel(); - } - for id in previous - .tool_runtime - .interrupt_for_run_replacement() - .await - { - let _ = handle.emit(&codec::abort(id)); - } - Some(previous.finished.clone()) - } else { - None - }; - let (results, result_receiver) = tool_result_channel(); - let (runtime_actions, runtime_action_receiver) = - mpsc::unbounded_channel::(); - let tool_runtime = tool_runtime_factory.next_run(); - let tools = ToolDispatcher::with_results( - tool_runtime.clone(), - results.clone(), - dependencies.store.clone(), - dependencies.web_cache.clone(), - ); - let generation = RunGeneration { - superseded: CancellationToken::new(), - finished: CancellationToken::new(), - run: Arc::new(parking_lot::Mutex::new(None)), - results, - runtime_actions, - tool_runtime, - tools, - }; - current = Some(generation.clone()); - spawn_run_request( - registry.clone(), - handle.clone(), + start_generation( + ®istry, + &handle, + &dependencies, + &blob_sync, + &context_sync, + &tool_runtime_factory, + &mut current, + &mut next_generation, request, - dependencies.clone(), - blob_sync.clone(), - context_sync.clone(), - generation, - previous_finished, - result_receiver, - runtime_action_receiver, - ); + ) + .await; } Some(pb::agent_client_message::Message::ExecClientMessage( message, @@ -331,27 +392,48 @@ impl ConversationRuntime { // return an explicit Protocol Error rather than falling through silently. Some( pb::agent_client_message::Message::ConversationAction( - action, + conversation_action, ), - ) => match action.action { + ) => match conversation_action.action.clone() { Some( pb::conversation_action::Action::UserMessageAction( action, ), ) => { - let Some(generation) = current.as_ref() else { + let delivered_to_active_run = + current.as_ref().is_some_and(|generation| { + generation.run.lock().as_ref().is_some_and( + |run| run.phase() == RunPhase::Running, + ) && generation + .runtime_actions + .send(compile::RuntimeAction::UserMessage( + action.clone(), + )) + .is_ok() + }); + if delivered_to_active_run { + continue; + } + let Some(previous) = current.as_ref() else { continue; }; - if generation - .runtime_actions - .send(compile::RuntimeAction::UserMessage(action)) - .is_err() - { - generation.results.send_error(crate::Error::Protocol( - "UserMessageAction arrived without an active Run" - .into(), - )); - } + let mut request = previous.request.clone(); + request.action = Some(conversation_action); + request.conversation_state = None; + request.pre_fetched_blobs.clear(); + waiting_for_action = false; + start_generation( + ®istry, + &handle, + &dependencies, + &blob_sync, + &context_sync, + &tool_runtime_factory, + &mut current, + &mut next_generation, + request, + ) + .await; } Some(pb::conversation_action::Action::CancelAction(_)) => { if let Some(generation) = current.as_ref() { @@ -428,6 +510,95 @@ impl ConversationRuntime { } } +#[allow(clippy::too_many_arguments)] +async fn start_generation( + registry: &ConversationRegistry, + handle: &TransportHandle, + dependencies: &ConversationDependencies, + blob_sync: &BlobSynchronizer, + context_sync: &RequestContextSynchronizer, + tool_runtime_factory: &CursorToolRuntime, + current: &mut Option, + next_generation: &mut u64, + request: pb::AgentRunRequest, +) { + let previous_finished = if let Some(previous) = current.take() { + previous.superseded.cancel(); + if let Some(run) = previous.run.lock().clone() { + run.cancel(); + } + for id in previous.tool_runtime.interrupt_for_run_replacement().await { + let _ = handle.emit(&codec::abort(id)); + } + Some(previous.finished.clone()) + } else { + None + }; + let (results, result_receiver) = tool_result_channel(); + let (runtime_actions, runtime_action_receiver) = + mpsc::unbounded_channel::(); + let tool_runtime = tool_runtime_factory.next_run(); + let tools = ToolDispatcher::with_results( + tool_runtime.clone(), + results.clone(), + dependencies.store.clone(), + dependencies.web_cache.clone(), + ); + let generation = RunGeneration { + id: *next_generation, + request: request.clone(), + superseded: CancellationToken::new(), + finished: CancellationToken::new(), + run: Arc::new(parking_lot::Mutex::new(None)), + results, + runtime_actions, + tool_runtime, + tools, + }; + *next_generation = next_generation.saturating_add(1); + *current = Some(generation.clone()); + spawn_run_request( + registry.clone(), + handle.clone(), + request, + dependencies.clone(), + blob_sync.clone(), + context_sync.clone(), + generation, + previous_finished, + result_receiver, + runtime_action_receiver, + ); +} + +fn finish_pending( + handle: &TransportHandle, + current: &Option, + pending: Option<(u64, TransportFinish)>, +) { + let Some((generation, finish)) = pending else { + return; + }; + if current + .as_ref() + .is_some_and(|current| current.id == generation) + { + finish_transport(handle, finish); + } +} + +fn finish_transport(handle: &TransportHandle, finish: TransportFinish) { + match finish { + TransportFinish::Success => super::finish_success(handle), + TransportFinish::Failed(error) => { + let _ = super::finish_failed(handle, &error); + } + TransportFinish::Cancelled => { + let _ = super::finish_cancelled(handle); + } + } +} + #[allow(clippy::too_many_arguments)] fn spawn_run_request( registry: ConversationRegistry, @@ -487,8 +658,12 @@ fn spawn_run_request( %error, "failed to prepare Cursor Run" ); - let _ = super::finish_failed(&handle, &error); - let _ = handle.command(TransportCommand::Close).await; + let _ = handle + .command(TransportCommand::RunFinished { + generation: generation.id, + finish: RunFinish::Transport(TransportFinish::Failed(error)), + }) + .await; return; } }; @@ -520,8 +695,12 @@ fn spawn_run_request( { CommandResult::Applied | CommandResult::Duplicate => { if !generation.superseded.is_cancelled() { - super::finish_success(&handle); - let _ = handle.command(TransportCommand::Close).await; + let _ = handle + .command(TransportCommand::RunFinished { + generation: generation.id, + finish: RunFinish::Transport(TransportFinish::Success), + }) + .await; } return; } @@ -545,8 +724,12 @@ fn spawn_run_request( } CommandResult::StaleTarget => { if !generation.superseded.is_cancelled() { - super::finish_success(&handle); - let _ = handle.command(TransportCommand::Close).await; + let _ = handle + .command(TransportCommand::RunFinished { + generation: generation.id, + finish: RunFinish::Transport(TransportFinish::Success), + }) + .await; } return; } @@ -605,16 +788,21 @@ fn spawn_run_request( tool_runtime: generation.tool_runtime.clone(), }, ); - if let Err(error) = output.run().await { - if !generation.superseded.is_cancelled() { - tracing::error!( - request_id = handle.request_id(), - %error, - "Cursor session failed" - ); - let _ = super::finish_failed(&handle, &error); + let finish = match output.run().await { + Ok(finish) => finish, + Err(error) => { + if generation.superseded.is_cancelled() { + RunFinish::Transport(TransportFinish::Cancelled) + } else { + tracing::error!( + request_id = handle.request_id(), + %error, + "Cursor session failed" + ); + RunFinish::Transport(TransportFinish::Failed(error)) + } } - } + }; let _ = core_run.await; registry.release(&conversation_id, &run_id).await; if generation @@ -626,7 +814,12 @@ fn spawn_run_request( *generation.run.lock() = None; } if !generation.superseded.is_cancelled() { - let _ = handle.command(TransportCommand::Close).await; + let _ = handle + .command(TransportCommand::RunFinished { + generation: generation.id, + finish, + }) + .await; } }); } diff --git a/server/src/cursor/services/blob_sync.rs b/server/src/cursor/services/blob_sync.rs index e88947a..316ce7d 100644 --- a/server/src/cursor/services/blob_sync.rs +++ b/server/src/cursor/services/blob_sync.rs @@ -76,22 +76,20 @@ impl BlobSynchronizer { let id = self.inner.store.put_blob(data, edges).await?; let result = self.ensure_set(&id, data).await; if let Some(trace) = self.inner.handle.trace() { - trace - .linked_blob( - "blob_set", - "byok_server", - &id, - serde_json::json!({ - "byte_count": data.len(), - "status": if result.is_ok() { "acknowledged" } else { "error" }, - "error": result.as_ref().err().map(ToString::to_string), - "edges": edges.iter().map(|edge| serde_json::json!({ - "child_blob_id": edge.child.to_base64(), - "field_name": edge.field_name, - })).collect::>(), - }), - ) - .await; + trace.linked_blob( + "blob_set", + "byok_server", + &id, + serde_json::json!({ + "byte_count": data.len(), + "status": if result.is_ok() { "acknowledged" } else { "error" }, + "error": result.as_ref().err().map(ToString::to_string), + "edges": edges.iter().map(|edge| serde_json::json!({ + "child_blob_id": edge.child.to_base64(), + "field_name": edge.field_name, + })).collect::>(), + }), + ); } result?; Ok(id) @@ -144,18 +142,16 @@ impl BlobSynchronizer { pub async fn get(&self, blob_id: &BlobId) -> Result>> { if let Some(data) = self.inner.store.get_blob(blob_id).await? { if let Some(trace) = self.inner.handle.trace() { - trace - .linked_blob( - "blob_get", - "byok_server", - blob_id, - serde_json::json!({ - "byte_count": data.len(), - "source": "local_store", - "status": "found", - }), - ) - .await; + trace.linked_blob( + "blob_get", + "byok_server", + blob_id, + serde_json::json!({ + "byte_count": data.len(), + "source": "local_store", + "status": "found", + }), + ); } return Ok(Some(data)); } @@ -194,45 +190,39 @@ impl BlobSynchronizer { if let Some(trace) = self.inner.handle.trace() { match &result { Ok(Some(data)) => { - trace - .linked_blob( - "blob_get", - "cursor_client", - blob_id, - serde_json::json!({ - "byte_count": data.len(), - "source": "cursor_client", - "status": "found", - }), - ) - .await; + trace.linked_blob( + "blob_get", + "cursor_client", + blob_id, + serde_json::json!({ + "byte_count": data.len(), + "source": "cursor_client", + "status": "found", + }), + ); } Ok(None) => { - trace - .artifact( - "blob_get", - "cursor_client", - &[], - serde_json::json!({ - "blob_id": blob_id.to_base64(), - "status": "missing", - }), - ) - .await; + trace.artifact( + "blob_get", + "cursor_client", + &[], + serde_json::json!({ + "blob_id": blob_id.to_base64(), + "status": "missing", + }), + ); } Err(error) => { - trace - .artifact( - "blob_get", - "cursor_client", - &[], - serde_json::json!({ - "blob_id": blob_id.to_base64(), - "status": "error", - "error": error.to_string(), - }), - ) - .await; + trace.artifact( + "blob_get", + "cursor_client", + &[], + serde_json::json!({ + "blob_id": blob_id.to_base64(), + "status": "error", + "error": error.to_string(), + }), + ); } } } diff --git a/server/src/cursor/services/observability.rs b/server/src/cursor/services/observability.rs deleted file mode 100644 index 56b9859..0000000 --- a/server/src/cursor/services/observability.rs +++ /dev/null @@ -1,225 +0,0 @@ -//! Records Cursor request traces and artifacts. -use std::{ - sync::{ - atomic::{AtomicBool, Ordering}, - Arc, - }, - time::{Duration, Instant}, -}; - -use tokio::sync::Mutex; - -use crate::store::{BlobId, BufferedCursorTraceChunk, Store}; - -#[derive(Clone)] -pub struct CursorTraceRecorder { - store: Store, - request_id: String, - chunks: Arc>, - finished: Arc, -} - -#[derive(Default)] -struct TraceChunkBuffer { - chunks: Vec, - bytes: usize, - first_chunk_at: Option, - generation: u64, -} - -const MAX_BUFFERED_CHUNKS: usize = 32; -const MAX_BUFFERED_BYTES: usize = 256 * 1024; -const MAX_BUFFER_AGE: Duration = Duration::from_millis(50); - -impl CursorTraceRecorder { - pub async fn begin( - store: Store, - request_id: &str, - conversation_id: Option<&str>, - route: &str, - model_id: Option<&str>, - ) -> Option { - match store - .start_cursor_trace_if_detailed(request_id, conversation_id, route, model_id) - .await - { - Ok(true) => Some(Self { - store, - request_id: request_id.into(), - chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())), - finished: Arc::new(AtomicBool::new(false)), - }), - Ok(false) => None, - Err(error) => { - tracing::warn!(request_id, %error, "failed to start Cursor trace"); - None - } - } - } - - pub async fn resume(store: Store, request_id: &str) -> Option { - match store.cursor_trace_exists(request_id).await { - Ok(true) => Some(Self { - store, - request_id: request_id.into(), - chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())), - finished: Arc::new(AtomicBool::new(false)), - }), - Ok(false) => None, - Err(error) => { - tracing::warn!(request_id, %error, "failed to resume Cursor trace"); - None - } - } - } - - pub fn request_id(&self) -> &str { - &self.request_id - } - - pub async fn request(&self, artifact_type: &str, data: &[u8], metadata: serde_json::Value) { - if let Err(error) = self - .store - .append_cursor_trace_artifact( - &self.request_id, - artifact_type, - "cursor_client", - data, - &metadata, - ) - .await - { - tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor request artifact"); - return; - } - if let Err(error) = self - .store - .add_cursor_trace_request_bytes(&self.request_id, data.len()) - .await - { - tracing::warn!(request_id = self.request_id, %error, "failed to update Cursor request trace size"); - } - } - - pub async fn artifact( - &self, - artifact_type: &str, - source: &str, - data: &[u8], - metadata: serde_json::Value, - ) { - if let Err(error) = self - .store - .append_cursor_trace_artifact(&self.request_id, artifact_type, source, data, &metadata) - .await - { - tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to record Cursor trace artifact"); - } - } - - pub async fn linked_blob( - &self, - artifact_type: &str, - source: &str, - blob_id: &BlobId, - metadata: serde_json::Value, - ) { - if let Err(error) = self - .store - .link_cursor_trace_artifact(&self.request_id, artifact_type, source, blob_id, &metadata) - .await - { - tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to link Cursor trace Blob"); - } - } - - pub async fn response_started(&self, status: u16) { - if let Err(error) = self - .store - .start_cursor_trace_response(&self.request_id, status) - .await - { - tracing::warn!(request_id = self.request_id, %error, "failed to start Cursor response trace"); - } - } - - pub async fn response_chunk(&self, source: &str, data: &[u8]) { - let mut buffer = self.chunks.lock().await; - if self.finished.load(Ordering::Acquire) { - return; - } - let schedule_flush = if buffer.chunks.is_empty() { - buffer.generation = buffer.generation.wrapping_add(1); - buffer.first_chunk_at = Some(Instant::now()); - Some(buffer.generation) - } else { - None - }; - buffer.bytes += data.len(); - buffer - .chunks - .push(BufferedCursorTraceChunk::new(source, data)); - let expired = buffer - .first_chunk_at - .is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE); - if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS - || buffer.bytes >= MAX_BUFFERED_BYTES - || expired - { - if let Err(error) = self.flush_locked(&mut buffer).await { - tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk"); - } - } - drop(buffer); - if let Some(generation) = schedule_flush { - let recorder = self.clone(); - tokio::spawn(async move { - tokio::time::sleep(MAX_BUFFER_AGE).await; - let mut buffer = recorder.chunks.lock().await; - if buffer.generation == generation { - if let Err(error) = recorder.flush_locked(&mut buffer).await { - tracing::warn!(request_id = recorder.request_id, %error, "failed to flush Cursor response chunks"); - } - } - }); - } - } - - pub async fn finish(&self, error: Option<&str>) { - if self.finished.swap(true, Ordering::AcqRel) { - return; - } - let mut buffer = self.chunks.lock().await; - if let Err(store_error) = self.flush_locked(&mut buffer).await { - tracing::warn!(request_id = self.request_id, %store_error, "failed to flush Cursor response chunks"); - } - drop(buffer); - if let Err(store_error) = self - .store - .finish_cursor_trace(&self.request_id, error) - .await - { - tracing::warn!(request_id = self.request_id, %store_error, "failed to finish Cursor trace"); - } - } - - async fn flush_locked(&self, buffer: &mut TraceChunkBuffer) -> crate::Result<()> { - if buffer.chunks.is_empty() { - return Ok(()); - } - let chunks = std::mem::take(&mut buffer.chunks); - buffer.bytes = 0; - buffer.first_chunk_at = None; - if let Err(error) = self - .store - .add_cursor_trace_response_chunks(&self.request_id, &chunks) - .await - { - buffer.bytes = chunks.iter().map(|chunk| chunk.data.len()).sum(); - buffer.first_chunk_at = Some(Instant::now()); - buffer.chunks = chunks; - return Err(error); - } - Ok(()) - } -} diff --git a/server/src/cursor/services/observability/event.rs b/server/src/cursor/services/observability/event.rs new file mode 100644 index 0000000..9491b7a --- /dev/null +++ b/server/src/cursor/services/observability/event.rs @@ -0,0 +1,71 @@ +use std::sync::{atomic::AtomicU8, Arc}; + +use bytes::Bytes; + +use crate::store::BlobId; + +pub(super) const TRACE_UNKNOWN: u8 = 0; +pub(super) const TRACE_ACTIVE: u8 = 1; +pub(super) const TRACE_DISABLED: u8 = 2; + +pub(super) enum TraceEvent { + Begin { + request_id: String, + activation: Arc, + conversation_id: Option, + route: String, + model_id: Option, + }, + Resume { + request_id: String, + activation: Arc, + }, + Request { + request_id: String, + artifact_type: String, + data: Bytes, + metadata: serde_json::Value, + }, + Artifact { + request_id: String, + artifact_type: String, + source: String, + data: Bytes, + metadata: serde_json::Value, + }, + LinkedBlob { + request_id: String, + artifact_type: String, + source: String, + blob_id: BlobId, + metadata: serde_json::Value, + }, + ResponseStarted { + request_id: String, + status: u16, + }, + ResponseChunk { + request_id: String, + source: String, + data: Bytes, + }, + Finish { + request_id: String, + error: Option, + }, +} + +impl TraceEvent { + pub(super) fn request_id(&self) -> &str { + match self { + Self::Begin { request_id, .. } + | Self::Resume { request_id, .. } + | Self::Request { request_id, .. } + | Self::Artifact { request_id, .. } + | Self::LinkedBlob { request_id, .. } + | Self::ResponseStarted { request_id, .. } + | Self::ResponseChunk { request_id, .. } + | Self::Finish { request_id, .. } => request_id, + } + } +} diff --git a/server/src/cursor/services/observability/mod.rs b/server/src/cursor/services/observability/mod.rs new file mode 100644 index 0000000..76727cf --- /dev/null +++ b/server/src/cursor/services/observability/mod.rs @@ -0,0 +1,157 @@ +//! Records Cursor request traces without blocking request or runtime paths. + +mod event; +mod worker; + +use std::sync::{ + atomic::{AtomicBool, AtomicU8, Ordering}, + Arc, +}; + +use bytes::Bytes; +use tokio::sync::mpsc; + +use crate::store::{BlobId, Store}; + +use event::{TraceEvent, TRACE_DISABLED, TRACE_UNKNOWN}; + +const TRACE_QUEUE_CAPACITY: usize = 512; + +#[derive(Clone)] +pub struct CursorTraceService { + sender: mpsc::Sender, +} + +impl CursorTraceService { + pub fn new(store: Store) -> Self { + let (sender, receiver) = mpsc::channel(TRACE_QUEUE_CAPACITY); + tokio::spawn(worker::run(store, receiver)); + Self { sender } + } + + pub fn recorder(&self, request_id: &str) -> CursorTraceRecorder { + CursorTraceRecorder { + request_id: Arc::from(request_id), + sender: self.sender.clone(), + finished: Arc::new(AtomicBool::new(false)), + activation: Arc::new(AtomicU8::new(TRACE_UNKNOWN)), + } + } +} + +#[derive(Clone)] +pub struct CursorTraceRecorder { + request_id: Arc, + sender: mpsc::Sender, + finished: Arc, + activation: Arc, +} + +impl CursorTraceRecorder { + pub fn request_id(&self) -> &str { + &self.request_id + } + + pub fn begin(&self, conversation_id: Option<&str>, route: &str, model_id: Option<&str>) { + self.send_control(TraceEvent::Begin { + request_id: self.request_id.to_string(), + activation: self.activation.clone(), + conversation_id: conversation_id.map(str::to_owned), + route: route.to_owned(), + model_id: model_id.map(str::to_owned), + }); + } + + pub fn resume(&self) { + self.send_control(TraceEvent::Resume { + request_id: self.request_id.to_string(), + activation: self.activation.clone(), + }); + } + + pub fn request(&self, artifact_type: &str, data: Bytes, metadata: serde_json::Value) { + self.send(TraceEvent::Request { + request_id: self.request_id.to_string(), + artifact_type: artifact_type.to_owned(), + data, + metadata, + }); + } + + pub fn artifact( + &self, + artifact_type: &str, + source: &str, + data: &[u8], + metadata: serde_json::Value, + ) { + self.send(TraceEvent::Artifact { + request_id: self.request_id.to_string(), + artifact_type: artifact_type.to_owned(), + source: source.to_owned(), + data: Bytes::copy_from_slice(data), + metadata, + }); + } + + pub fn linked_blob( + &self, + artifact_type: &str, + source: &str, + blob_id: &BlobId, + metadata: serde_json::Value, + ) { + self.send(TraceEvent::LinkedBlob { + request_id: self.request_id.to_string(), + artifact_type: artifact_type.to_owned(), + source: source.to_owned(), + blob_id: blob_id.clone(), + metadata, + }); + } + + pub fn response_started(&self, status: u16) { + self.send(TraceEvent::ResponseStarted { + request_id: self.request_id.to_string(), + status, + }); + } + + pub fn response_chunk(&self, source: &str, data: Bytes) { + if self.finished.load(Ordering::Acquire) { + return; + } + self.send(TraceEvent::ResponseChunk { + request_id: self.request_id.to_string(), + source: source.to_owned(), + data, + }); + } + + pub fn finish(&self, error: Option<&str>) { + if self.finished.swap(true, Ordering::AcqRel) { + return; + } + self.send_control(TraceEvent::Finish { + request_id: self.request_id.to_string(), + error: error.map(str::to_owned), + }); + } + + fn send(&self, event: TraceEvent) { + if self.activation.load(Ordering::Acquire) == TRACE_DISABLED { + return; + } + self.send_control(event); + } + + fn send_control(&self, event: TraceEvent) { + if let Err(error) = self.sender.try_send(event) { + tracing::warn!( + request_id = %self.request_id, + %error, + "dropping Cursor trace event" + ); + } + } +} diff --git a/server/src/cursor/services/observability/worker.rs b/server/src/cursor/services/observability/worker.rs new file mode 100644 index 0000000..62e6046 --- /dev/null +++ b/server/src/cursor/services/observability/worker.rs @@ -0,0 +1,349 @@ +use std::{ + collections::{BTreeMap, HashMap}, + sync::atomic::Ordering, + time::Duration, +}; + +use tokio::sync::mpsc; + +use crate::store::{BufferedCursorTraceChunk, Store}; + +use super::event::{TraceEvent, TRACE_ACTIVE, TRACE_DISABLED}; + +const MAX_BUFFERED_CHUNKS: usize = 32; +const MAX_BUFFERED_BYTES: usize = 256 * 1024; +const FLUSH_INTERVAL: Duration = Duration::from_millis(50); + +#[derive(Clone, Copy)] +enum TraceState { + Active, + Disabled, +} + +#[derive(Default)] +struct ResponseBuffer { + chunks: Vec, + bytes: usize, +} + +struct BufferedRequest { + artifact_type: String, + data: bytes::Bytes, + metadata: serde_json::Value, +} + +#[derive(Default)] +struct RequestOrder { + next: i64, + pending: BTreeMap>, +} + +pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver) { + let mut states = HashMap::::new(); + let mut buffers = HashMap::::new(); + let mut request_orders = HashMap::::new(); + let mut interval = tokio::time::interval(FLUSH_INTERVAL); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + loop { + tokio::select! { + event = receiver.recv() => { + let Some(event) = event else { + flush_all(&store, &mut buffers).await; + return; + }; + process(&store, &mut states, &mut buffers, &mut request_orders, event).await; + } + _ = interval.tick() => flush_all(&store, &mut buffers).await, + } + } +} + +async fn process( + store: &Store, + states: &mut HashMap, + buffers: &mut HashMap, + request_orders: &mut HashMap, + event: TraceEvent, +) { + let request_id = event.request_id().to_owned(); + let finishes_trace = matches!(&event, TraceEvent::Finish { .. }); + match event { + TraceEvent::Begin { + request_id, + activation, + conversation_id, + route, + model_id, + } => { + let state = match store + .start_cursor_trace_if_detailed( + &request_id, + conversation_id.as_deref(), + &route, + model_id.as_deref(), + ) + .await + { + Ok(true) => TraceState::Active, + Ok(false) => TraceState::Disabled, + Err(error) => { + tracing::warn!(%request_id, %error, "failed to start Cursor trace"); + TraceState::Disabled + } + }; + activation.store( + match state { + TraceState::Active => TRACE_ACTIVE, + TraceState::Disabled => TRACE_DISABLED, + }, + Ordering::Release, + ); + states.insert(request_id, state); + return; + } + TraceEvent::Resume { + request_id, + activation, + } => { + let state = ensure_state(store, states, &request_id).await; + activation.store( + match state { + TraceState::Active => TRACE_ACTIVE, + TraceState::Disabled => TRACE_DISABLED, + }, + Ordering::Release, + ); + return; + } + _ => {} + } + + if !matches!( + ensure_state(store, states, &request_id).await, + TraceState::Active + ) { + if finishes_trace { + states.remove(&request_id); + buffers.remove(&request_id); + request_orders.remove(&request_id); + } + return; + } + + let result = match event { + TraceEvent::Request { + artifact_type, + data, + metadata, + .. + } => { + append_request( + store, + request_orders, + &request_id, + artifact_type, + data, + metadata, + ) + .await + } + TraceEvent::Artifact { + artifact_type, + source, + data, + metadata, + .. + } => { + store + .append_cursor_trace_artifact( + &request_id, + &artifact_type, + &source, + &data, + &metadata, + ) + .await + } + TraceEvent::LinkedBlob { + artifact_type, + source, + blob_id, + metadata, + .. + } => { + store + .link_cursor_trace_artifact( + &request_id, + &artifact_type, + &source, + &blob_id, + &metadata, + ) + .await + } + TraceEvent::ResponseStarted { status, .. } => { + store.start_cursor_trace_response(&request_id, status).await + } + TraceEvent::ResponseChunk { source, data, .. } => { + let buffer = buffers.entry(request_id.clone()).or_default(); + buffer.bytes += data.len(); + buffer + .chunks + .push(BufferedCursorTraceChunk::new(&source, &data)); + if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS || buffer.bytes >= MAX_BUFFERED_BYTES { + flush_one(store, buffers, &request_id).await; + } + return; + } + TraceEvent::Finish { error, .. } => { + flush_request_order(store, request_orders, &request_id).await; + flush_one(store, buffers, &request_id).await; + store + .finish_cursor_trace(&request_id, error.as_deref()) + .await + } + TraceEvent::Begin { .. } | TraceEvent::Resume { .. } => unreachable!(), + }; + if let Err(error) = result { + tracing::warn!(%request_id, %error, "failed to record Cursor trace event"); + } + if finishes_trace { + states.remove(&request_id); + buffers.remove(&request_id); + } +} + +async fn append_request( + store: &Store, + request_orders: &mut HashMap, + request_id: &str, + artifact_type: String, + data: bytes::Bytes, + metadata: serde_json::Value, +) -> crate::Result<()> { + let append_seqno = metadata + .get("append_seqno") + .and_then(serde_json::Value::as_i64); + let ordered = artifact_type == "bidi_request" + && metadata + .get("accepted") + .and_then(serde_json::Value::as_bool) + == Some(true) + && metadata + .get("route_outcome") + .and_then(serde_json::Value::as_str) + == Some("local"); + let Some(append_seqno) = append_seqno.filter(|_| ordered) else { + return store + .append_cursor_trace_request( + request_id, + &artifact_type, + "cursor_client", + &data, + &metadata, + ) + .await; + }; + + let request = BufferedRequest { + artifact_type, + data, + metadata, + }; + let order = request_orders.entry(request_id.to_owned()).or_default(); + if append_seqno < order.next { + return store + .append_cursor_trace_request( + request_id, + &request.artifact_type, + "cursor_client", + &request.data, + &request.metadata, + ) + .await; + } + order.pending.entry(append_seqno).or_default().push(request); + while let Some(requests) = order.pending.remove(&order.next) { + for request in requests { + store + .append_cursor_trace_request( + request_id, + &request.artifact_type, + "cursor_client", + &request.data, + &request.metadata, + ) + .await?; + } + order.next = order.next.saturating_add(1); + } + Ok(()) +} + +async fn flush_request_order( + store: &Store, + request_orders: &mut HashMap, + request_id: &str, +) { + let Some(order) = request_orders.remove(request_id) else { + return; + }; + for requests in order.pending.into_values() { + for request in requests { + if let Err(error) = store + .append_cursor_trace_request( + request_id, + &request.artifact_type, + "cursor_client", + &request.data, + &request.metadata, + ) + .await + { + tracing::warn!(%request_id, %error, "failed to flush ordered Cursor request trace"); + } + } + } +} + +async fn ensure_state( + store: &Store, + states: &mut HashMap, + request_id: &str, +) -> TraceState { + if let Some(state) = states.get(request_id).copied() { + return state; + } + let state = match store.cursor_trace_exists(request_id).await { + Ok(true) => TraceState::Active, + Ok(false) => TraceState::Disabled, + Err(error) => { + tracing::warn!(%request_id, %error, "failed to resume Cursor trace"); + TraceState::Disabled + } + }; + states.insert(request_id.to_owned(), state); + state +} + +async fn flush_one(store: &Store, buffers: &mut HashMap, request_id: &str) { + let Some(mut buffer) = buffers.remove(request_id) else { + return; + }; + if let Err(error) = store + .add_cursor_trace_response_chunks(request_id, &buffer.chunks) + .await + { + tracing::warn!(%request_id, %error, "failed to flush Cursor response chunks"); + buffer.bytes = buffer.chunks.iter().map(|chunk| chunk.data.len()).sum(); + buffers.insert(request_id.to_owned(), buffer); + } +} + +async fn flush_all(store: &Store, buffers: &mut HashMap) { + let request_ids = buffers.keys().cloned().collect::>(); + for request_id in request_ids { + flush_one(store, buffers, &request_id).await; + } +} diff --git a/server/src/cursor/transport/handle.rs b/server/src/cursor/transport/handle.rs index 4b72b5b..87a0a2e 100644 --- a/server/src/cursor/transport/handle.rs +++ b/server/src/cursor/transport/handle.rs @@ -15,7 +15,7 @@ use crate::{ Error, Result, }; -use super::OutputHub; +use super::{OutputHub, TransportAdmission, TransportLifecycle}; #[derive(Clone, Debug, PartialEq, Eq)] pub struct TransportParent { @@ -30,7 +30,8 @@ pub struct TransportHandle { output: Arc, conversation_id: Arc>, parent: Arc>, - trace: Option, + trace: CursorTraceRecorder, + lifecycle: TransportLifecycle, disconnect: CancellationToken, } @@ -39,7 +40,7 @@ impl TransportHandle { request_id: String, commands: mpsc::Sender, output: Arc, - trace: Option, + trace: CursorTraceRecorder, ) -> Self { Self { request_id, @@ -48,6 +49,7 @@ impl TransportHandle { conversation_id: Arc::new(OnceLock::new()), parent: Arc::new(OnceLock::new()), trace, + lifecycle: TransportLifecycle::new(), disconnect: CancellationToken::new(), } } @@ -127,12 +129,46 @@ impl TransportHandle { self.output.close() } - pub(crate) async fn wait_closed(&self) { - self.output.wait_closed().await; + pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> { + Some(&self.trace) } - pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> { - self.trace.as_ref() + pub(crate) fn accepting_appends(&self) -> bool { + self.lifecycle.is_open() + } + + pub(crate) fn admit(&self) -> Result { + self.lifecycle + .admit() + .ok_or_else(|| Error::RunNotFound(self.request_id.clone())) + } + + pub(crate) fn begin_close(&self) { + self.lifecycle.begin_close(); + } + + pub(crate) fn admissions_drained(&self) -> bool { + self.lifecycle.admissions_drained() + } + + pub(crate) async fn wait_admissions_drained(&self) { + self.lifecycle.wait_admissions_drained().await; + } + + pub(crate) fn mark_draining(&self) { + self.lifecycle.mark_draining(); + } + + pub(crate) fn reopen(&self) { + self.lifecycle.reopen(); + } + + pub(crate) fn close_transport(&self) { + self.lifecycle.close(); + } + + pub(crate) async fn wait_transport_closed(&self) { + self.lifecycle.wait_closed().await; } pub(crate) fn disconnect_token(&self) -> CancellationToken { diff --git a/server/src/cursor/transport/lifecycle.rs b/server/src/cursor/transport/lifecycle.rs new file mode 100644 index 0000000..b1faeeb --- /dev/null +++ b/server/src/cursor/transport/lifecycle.rs @@ -0,0 +1,170 @@ +//! Coordinates append admission with transport shutdown. + +use std::sync::Arc; + +use tokio::sync::Notify; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum TransportState { + Open, + Closing, + Draining, + Closed, +} + +#[derive(Clone)] +pub(crate) struct TransportLifecycle { + inner: Arc, +} + +struct LifecycleInner { + state: parking_lot::Mutex, + admissions_drained: Notify, + closed: Notify, +} + +struct LifecycleState { + state: TransportState, + admissions: usize, +} + +pub(crate) struct TransportAdmission { + inner: Arc, +} + +impl TransportLifecycle { + pub(crate) fn new() -> Self { + Self { + inner: Arc::new(LifecycleInner { + state: parking_lot::Mutex::new(LifecycleState { + state: TransportState::Open, + admissions: 0, + }), + admissions_drained: Notify::new(), + closed: Notify::new(), + }), + } + } + + pub(crate) fn is_open(&self) -> bool { + self.inner.state.lock().state == TransportState::Open + } + + pub(crate) fn admit(&self) -> Option { + let mut lifecycle = self.inner.state.lock(); + if lifecycle.state != TransportState::Open { + return None; + } + lifecycle.admissions += 1; + Some(TransportAdmission { + inner: self.inner.clone(), + }) + } + + pub(crate) fn begin_close(&self) { + let mut lifecycle = self.inner.state.lock(); + if lifecycle.state != TransportState::Open { + return; + } + lifecycle.state = TransportState::Closing; + let drained = lifecycle.admissions == 0; + drop(lifecycle); + if drained { + self.inner.admissions_drained.notify_waiters(); + } + } + + pub(crate) fn admissions_drained(&self) -> bool { + self.inner.state.lock().admissions == 0 + } + + pub(crate) async fn wait_admissions_drained(&self) { + loop { + let notified = self.inner.admissions_drained.notified(); + if self.inner.state.lock().admissions == 0 { + return; + } + notified.await; + } + } + + pub(crate) fn mark_draining(&self) { + let mut lifecycle = self.inner.state.lock(); + if lifecycle.state == TransportState::Closing && lifecycle.admissions == 0 { + lifecycle.state = TransportState::Draining; + } + } + + pub(crate) fn reopen(&self) { + let mut lifecycle = self.inner.state.lock(); + if matches!( + lifecycle.state, + TransportState::Closing | TransportState::Draining + ) { + lifecycle.state = TransportState::Open; + } + } + + pub(crate) fn close(&self) { + let mut lifecycle = self.inner.state.lock(); + if lifecycle.state == TransportState::Closed { + return; + } + lifecycle.state = TransportState::Closed; + drop(lifecycle); + self.inner.closed.notify_waiters(); + } + + pub(crate) async fn wait_closed(&self) { + loop { + let notified = self.inner.closed.notified(); + if self.inner.state.lock().state == TransportState::Closed { + return; + } + notified.await; + } + } +} + +impl Drop for TransportAdmission { + fn drop(&mut self) { + let mut lifecycle = self.inner.state.lock(); + lifecycle.admissions = lifecycle.admissions.saturating_sub(1); + let drained = lifecycle.state == TransportState::Closing && lifecycle.admissions == 0; + drop(lifecycle); + if drained { + self.inner.admissions_drained.notify_waiters(); + } + } +} + +#[cfg(test)] +mod tests { + use super::{TransportLifecycle, TransportState}; + + #[tokio::test] + async fn closing_waits_for_existing_admissions() { + let lifecycle = TransportLifecycle::new(); + let admission = lifecycle.admit().unwrap(); + lifecycle.begin_close(); + assert!(lifecycle.admit().is_none()); + drop(admission); + lifecycle.wait_admissions_drained().await; + lifecycle.mark_draining(); + assert_eq!(lifecycle.inner.state.lock().state, TransportState::Draining); + } + + #[tokio::test] + async fn an_admitted_continuation_reopens_the_transport() { + let lifecycle = TransportLifecycle::new(); + let admission = lifecycle.admit().unwrap(); + lifecycle.begin_close(); + drop(admission); + lifecycle.wait_admissions_drained().await; + lifecycle.mark_draining(); + lifecycle.reopen(); + + assert_eq!(lifecycle.inner.state.lock().state, TransportState::Open); + assert!(lifecycle.admit().is_some()); + } +} diff --git a/server/src/cursor/transport/mod.rs b/server/src/cursor/transport/mod.rs index 20815bd..263e02c 100644 --- a/server/src/cursor/transport/mod.rs +++ b/server/src/cursor/transport/mod.rs @@ -2,10 +2,12 @@ mod handle; mod inbox; +mod lifecycle; mod output; mod registry; pub use handle::*; pub use inbox::*; +pub(crate) use lifecycle::*; pub use output::*; pub use registry::*; diff --git a/server/src/cursor/transport/output.rs b/server/src/cursor/transport/output.rs index 0f513d9..bbd55ef 100644 --- a/server/src/cursor/transport/output.rs +++ b/server/src/cursor/transport/output.rs @@ -1,12 +1,11 @@ //! Buffers, replays, broadcasts, and atomically closes downstream output. use bytes::Bytes; -use tokio::sync::{mpsc, Notify}; +use tokio::sync::mpsc; #[derive(Default)] pub struct OutputHub { state: parking_lot::Mutex, - closed: Notify, } #[derive(Default)] @@ -49,17 +48,6 @@ impl OutputHub { state.closed = true; state.subscribers.clear(); drop(state); - self.closed.notify_waiters(); true } - - pub async fn wait_closed(&self) { - loop { - let notified = self.closed.notified(); - if self.state.lock().closed { - return; - } - notified.await; - } - } } diff --git a/server/src/cursor/transport/registry.rs b/server/src/cursor/transport/registry.rs index db57032..687987c 100644 --- a/server/src/cursor/transport/registry.rs +++ b/server/src/cursor/transport/registry.rs @@ -1,13 +1,19 @@ //! Maps request IDs to active transport handles. -use std::{collections::HashMap, sync::Arc}; +use std::{ + collections::HashMap, + sync::{ + atomic::{AtomicU64, Ordering}, + Arc, + }, +}; use tokio::sync::{mpsc, Mutex, Notify}; use crate::{ cursor::{ conversation::ConversationRegistry, prompting::PromptCompiler, - services::observability::CursorTraceRecorder, + services::observability::CursorTraceService, }, plugin::PluginRegistry, provider::Provider, @@ -24,15 +30,23 @@ pub struct TransportRegistry { } struct RegistryInner { - local: Mutex>, + local: Mutex>, + next_local_generation: AtomicU64, upstream: Mutex>, route_changed: Notify, store: Store, + traces: CursorTraceService, web_cache: WebCache, plugins: Option, conversations: ConversationRegistry, } +#[derive(Clone)] +struct LocalTransport { + generation: u64, + handle: TransportHandle, +} + #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum TransportRoute { Local, @@ -99,8 +113,10 @@ impl TransportRegistry { Self { inner: Arc::new(RegistryInner { local: Mutex::new(HashMap::new()), + next_local_generation: AtomicU64::new(1), upstream: Mutex::new(HashMap::new()), route_changed: Notify::new(), + traces: CursorTraceService::new(store.clone()), conversations: ConversationRegistry::new( store.clone(), provider, @@ -119,6 +135,13 @@ impl TransportRegistry { &self.inner.store } + pub fn trace( + &self, + request_id: &str, + ) -> crate::cursor::services::observability::CursorTraceRecorder { + self.inner.traces.recorder(request_id) + } + pub fn web_cache(&self) -> &WebCache { &self.inner.web_cache } @@ -132,18 +155,37 @@ impl TransportRegistry { } pub async fn get_or_create(&self, request_id: &str) -> Result { - if let Some(handle) = self.inner.local.lock().await.get(request_id).cloned() { - return Ok(handle); + self.get_or_create_for_append(request_id, false).await + } + + pub(crate) async fn get_or_create_for_append( + &self, + request_id: &str, + replace_closing: bool, + ) -> Result { + let mut local = self.inner.local.lock().await; + if let Some(transport) = local.get(request_id) { + if transport.handle.accepting_appends() || !replace_closing { + return Ok(transport.handle.clone()); + } } + local.remove(request_id); let (commands, receiver) = mpsc::channel(128); let output = Arc::new(OutputHub::default()); - let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await; - let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace); - let mut local = self.inner.local.lock().await; - if let Some(existing) = local.get(request_id).cloned() { - return Ok(existing); - } - local.insert(request_id.into(), handle.clone()); + let trace = self.inner.traces.recorder(request_id); + trace.resume(); + let handle = TransportHandle::new(request_id.into(), commands, output, trace); + let generation = self + .inner + .next_local_generation + .fetch_add(1, Ordering::Relaxed); + local.insert( + request_id.into(), + LocalTransport { + generation, + handle: handle.clone(), + }, + ); drop(local); self.inner.route_changed.notify_waiters(); self.inner @@ -152,17 +194,29 @@ impl TransportRegistry { let registry = Arc::downgrade(&self.inner); let request_id = request_id.to_string(); + let lifecycle = handle.clone(); tokio::spawn(async move { - output.wait_closed().await; + lifecycle.wait_transport_closed().await; if let Some(registry) = registry.upgrade() { - registry.local.lock().await.remove(&request_id); + let mut local = registry.local.lock().await; + if local + .get(&request_id) + .is_some_and(|transport| transport.generation == generation) + { + local.remove(&request_id); + } } }); Ok(handle) } pub async fn local(&self, request_id: &str) -> Option { - self.inner.local.lock().await.get(request_id).cloned() + self.inner + .local + .lock() + .await + .get(request_id) + .map(|transport| transport.handle.clone()) } pub async fn mark_upstream(&self, request_id: &str) { @@ -206,10 +260,13 @@ impl TransportRegistry { self.inner.conversations.shutdown().await; let handles = std::mem::take(&mut *self.inner.local.lock().await); self.inner.upstream.lock().await.clear(); - for handle in handles.into_values() { - handle.disconnect().await; - let _ = - tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await; + for transport in handles.into_values() { + transport.handle.disconnect().await; + let _ = tokio::time::timeout( + std::time::Duration::from_secs(2), + transport.handle.wait_transport_closed(), + ) + .await; } } } diff --git a/server/src/model/projection.rs b/server/src/model/projection.rs index 05dfaee..c6a7774 100644 --- a/server/src/model/projection.rs +++ b/server/src/model/projection.rs @@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize}; use crate::{Error, Result}; use super::{ - CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent, - ToolResultContent, + normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, + ToolCallContent, ToolResultContent, }; #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] @@ -91,7 +91,7 @@ fn project_tool_round( "tool round repeats provider replay state".into(), )); } - calls.extend(part_calls.iter().cloned()); + calls.extend(part_calls.iter().map(normalized_tool_call)); cursor += 1; while cursor < messages.len() { @@ -139,7 +139,7 @@ fn project_tool_round( .map(|(message_id, result)| ProjectedMessage { message_id, role: Role::Tool, - content: ProjectedContent::ToolResult(result), + content: ProjectedContent::ToolResult(normalized_tool_result(&result)), }), ); Ok(Some((output, cursor))) @@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage { text: text.clone(), thinking: thinking.clone(), replay_state: replay_state.clone(), - calls: tool_calls.clone(), + calls: tool_calls.iter().map(normalized_tool_call).collect(), }, - MessageContent::ToolResult(result) => ProjectedContent::ToolResult(result.clone()), + MessageContent::ToolResult(result) => { + ProjectedContent::ToolResult(normalized_tool_result(result)) + } }; ProjectedMessage { message_id: message.message_id.clone(), @@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage { content, } } + +fn normalized_tool_call(call: &ToolCallContent) -> ToolCallContent { + let mut call = call.clone(); + call.name = normalize_tool_name(&call.name); + call +} + +fn normalized_tool_result(result: &ToolResultContent) -> ToolResultContent { + let mut result = result.clone(); + result.name = normalize_tool_name(&result.name); + result +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::Origin; + use serde_json::json; + + #[test] + fn tool_names_are_normalized_before_provider_dispatch() { + let messages = [CanonicalMessage { + message_id: "assistant-1".into(), + role: Role::Assistant, + origin: Origin::Assistant, + content: MessageContent::Assistant { + text: String::new(), + thinking: String::new(), + tool_round_id: None, + replay_state: None, + tool_calls: vec![ToolCallContent { + index: 0, + call_id: "call-1".into(), + name: "multi_tool_use.parallel".into(), + arguments: json!({}), + }], + }, + runtime_event_id: None, + }]; + + let projected = project_messages(&messages).unwrap(); + let ProjectedContent::Assistant { calls, .. } = &projected[0].content else { + panic!("expected assistant projection"); + }; + assert_eq!(calls[0].name, "multi_tool_use_parallel"); + } +} diff --git a/server/src/model/tool.rs b/server/src/model/tool.rs index c522765..c35f115 100644 --- a/server/src/model/tool.rs +++ b/server/src/model/tool.rs @@ -4,6 +4,24 @@ use serde_json::Value; use super::ProviderReplayState; +pub fn normalize_tool_name(name: &str) -> String { + let normalized = name + .chars() + .map(|character| { + if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') { + character + } else { + '_' + } + }) + .collect::(); + if normalized.is_empty() { + "_".into() + } else { + normalized + } +} + #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] pub struct ToolDefinition { pub name: String, diff --git a/server/src/run/model_cycle.rs b/server/src/run/model_cycle.rs index e6a24ba..f9d3b7a 100644 --- a/server/src/run/model_cycle.rs +++ b/server/src/run/model_cycle.rs @@ -6,7 +6,7 @@ use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; use crate::{ - model::{ProviderReplayState, ToolCall, Usage}, + model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage}, provider::{FinishReason, ModelEvent, ProviderStream}, }; @@ -165,6 +165,7 @@ pub async fn consume_model_cycle( call_id, name, } => { + let name = normalize_tool_name(&name); let Some(model_call_id) = model_call_id.as_ref() else { return Err(failure( RunFailure::Protocol("provider emitted content before Start".into()), @@ -406,6 +407,34 @@ mod tests { }; use tokio_stream::wrappers::ReceiverStream; + #[tokio::test] + async fn provider_tool_names_are_normalized_when_received() { + let events = vec![ + Ok(ModelEvent::Start { + model_call_id: "call".into(), + }), + Ok(ModelEvent::ToolCallStart { + index: 0, + call_id: "tool-call".into(), + name: "multi_tool_use.parallel".into(), + }), + Ok(ModelEvent::ToolCallEnd { index: 0 }), + Ok(ModelEvent::Done(FinishReason::ToolUse)), + ]; + let stream = Box::pin(tokio_stream::iter(events)); + let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4); + + let result = consume_model_cycle(stream, &event_tx, &CancellationToken::new()) + .await + .unwrap(); + + assert_eq!(result.calls[0].name, "multi_tool_use_parallel"); + assert!(matches!( + event_rx.recv().await, + Some(RunEvent::ToolCallStart { name, .. }) if name == "multi_tool_use_parallel" + )); + } + #[tokio::test] async fn usage_is_forwarded_before_the_provider_call_finishes() { let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4); diff --git a/server/src/store/cursor_traces.rs b/server/src/store/cursor_traces.rs index e91b2fe..55f8a33 100644 --- a/server/src/store/cursor_traces.rs +++ b/server/src/store/cursor_traces.rs @@ -62,6 +62,40 @@ impl Store { .await?) } + pub async fn append_cursor_trace_request( + &self, + request_id: &str, + artifact_type: &str, + source: &str, + data: &[u8], + metadata: &serde_json::Value, + ) -> Result<()> { + let metadata_json = serde_json::to_string(metadata)?; + let blob_id = BlobId::digest(data); + let _write = self.writes.lock().await; + let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + Self::put_blob_tx(&mut tx, &blob_id, data, &[]).await?; + Self::link_cursor_trace_artifact_tx( + &mut tx, + request_id, + artifact_type, + source, + &blob_id, + &metadata_json, + ) + .await?; + sqlx::query( + "UPDATE cursor_run_traces + SET request_bytes = request_bytes + ? WHERE request_id = ?", + ) + .bind(as_i64(data.len())) + .bind(request_id) + .execute(&mut *tx) + .await?; + tx.commit().await?; + Ok(()) + } + pub async fn append_cursor_trace_artifact( &self, request_id: &str, @@ -144,23 +178,6 @@ impl Store { Ok(()) } - pub async fn add_cursor_trace_request_bytes( - &self, - request_id: &str, - bytes: usize, - ) -> Result<()> { - let _write = self.writes.lock().await; - sqlx::query( - "UPDATE cursor_run_traces - SET request_bytes = request_bytes + ? WHERE request_id = ?", - ) - .bind(as_i64(bytes)) - .bind(request_id) - .execute(&self.pool) - .await?; - Ok(()) - } - pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> { let now = now_ms(); let _write = self.writes.lock().await; diff --git a/server/tests/cursor_trace_queue.rs b/server/tests/cursor_trace_queue.rs new file mode 100644 index 0000000..7debc58 --- /dev/null +++ b/server/tests/cursor_trace_queue.rs @@ -0,0 +1,129 @@ +//! Verifies that Cursor trace persistence is ordered and detached from producers. + +#[path = "support/fixtures.rs"] +mod fixtures; + +use std::time::{Duration, Instant}; + +use bytes::Bytes; +use cursor_server::{cursor::services::observability::CursorTraceService, store::Store}; +use sqlx::{Connection, SqliteConnection}; + +#[tokio::test] +async fn trace_producers_do_not_wait_for_sqlite_and_artifacts_stay_ordered() { + let directory = tempfile::tempdir().unwrap(); + let url = format!("sqlite://{}", directory.path().join("test.db").display()); + let store = Store::connect(&url).await.unwrap(); + store.set_detailed_logging(true).await.unwrap(); + let traces = CursorTraceService::new(store.clone()); + let recorder = traces.recorder("trace-queue-order"); + recorder.begin(Some("conversation-1"), "local_byok", Some("model-1")); + + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if store + .cursor_trace("trace-queue-order") + .await + .unwrap() + .is_some() + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + let mut write_lock = SqliteConnection::connect(&url).await.unwrap(); + sqlx::query("BEGIN IMMEDIATE") + .execute(&mut write_lock) + .await + .unwrap(); + recorder.request( + "bidi_request", + Bytes::from_static(b"request-0"), + serde_json::json!({ + "append_seqno": 0, + "accepted": true, + "route_outcome": "local" + }), + ); + tokio::time::sleep(Duration::from_millis(25)).await; + + let started = Instant::now(); + let mut seqnos = (1..64).collect::>(); + for pair in seqnos.chunks_mut(2) { + pair.reverse(); + } + for seqno in seqnos { + recorder.request( + "bidi_request", + Bytes::from(format!("request-{seqno}")), + serde_json::json!({ + "append_seqno": seqno, + "accepted": true, + "route_outcome": "local" + }), + ); + } + recorder.finish(None); + assert!(started.elapsed() < Duration::from_millis(100)); + sqlx::query("ROLLBACK") + .execute(&mut write_lock) + .await + .unwrap(); + + let artifacts = tokio::time::timeout(Duration::from_secs(5), async { + loop { + let artifacts = store + .cursor_trace_artifacts("trace-queue-order") + .await + .unwrap(); + if artifacts.len() == 64 { + break artifacts; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + for (expected, artifact) in artifacts.iter().enumerate() { + assert_eq!(artifact.seq, expected as i64); + assert_eq!(artifact.metadata["append_seqno"], expected as i64); + } + let trace = store + .cursor_trace("trace-queue-order") + .await + .unwrap() + .unwrap(); + assert_eq!(trace.status, "completed"); + assert_eq!( + trace.request_bytes, + (0..64) + .map(|seqno| format!("request-{seqno}").len() as i64) + .sum::() + ); +} + +#[tokio::test] +async fn events_for_disabled_detailed_logging_are_discarded_off_path() { + let (_directory, store) = fixtures::temp_store().await; + let traces = CursorTraceService::new(store.clone()); + let recorder = traces.recorder("trace-disabled"); + + recorder.begin(None, "local_byok", Some("model-1")); + recorder.request( + "bidi_request", + Bytes::from_static(b"body"), + serde_json::json!({"append_seqno": 0}), + ); + + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(store + .cursor_trace("trace-disabled") + .await + .unwrap() + .is_none()); +} diff --git a/server/tests/cursor_transport_lifecycle.rs b/server/tests/cursor_transport_lifecycle.rs new file mode 100644 index 0000000..258ed6a --- /dev/null +++ b/server/tests/cursor_transport_lifecycle.rs @@ -0,0 +1,72 @@ +//! Verifies registry ownership follows the transport actor rather than output subscriptions. + +#[path = "support/fake_provider.rs"] +mod fake_provider; +#[path = "support/fixtures.rs"] +mod fixtures; + +use std::{sync::Arc, time::Duration}; + +use cursor_server::cursor::{ + conversation::TransportCommand, + prompting::{PromptAssets, PromptCompiler}, + transport::TransportRegistry, +}; + +async fn registry() -> (tempfile::TempDir, TransportRegistry) { + 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(); + ( + directory, + TransportRegistry::new( + store, + Arc::new(fake_provider::FakeProvider::default()), + PromptCompiler::new(assets), + ), + ) +} + +#[tokio::test] +async fn actor_exit_removes_the_matching_transport_and_allows_a_new_generation() { + let (_directory, registry) = registry().await; + let first = registry.get_or_create("lifecycle-request").await.unwrap(); + assert!(registry.local("lifecycle-request").await.is_some()); + + first.command(TransportCommand::Disconnect).await.unwrap(); + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if registry.local("lifecycle-request").await.is_none() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + let second = registry.get_or_create("lifecycle-request").await.unwrap(); + assert_eq!(second.request_id(), "lifecycle-request"); + assert!(registry.local("lifecycle-request").await.is_some()); + second.command(TransportCommand::Disconnect).await.unwrap(); +} + +#[tokio::test] +async fn dropping_an_output_subscription_does_not_remove_the_transport() { + let (_directory, registry) = registry().await; + let handle = registry + .get_or_create("subscription-request") + .await + .unwrap(); + let subscription = handle.subscribe(); + drop(subscription); + + tokio::time::sleep(Duration::from_millis(25)).await; + assert!(registry.local("subscription-request").await.is_some()); + + handle.command(TransportCommand::Disconnect).await.unwrap(); +} diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 1cfd882..9c50c8f 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -422,6 +422,85 @@ async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() { assert_eq!(output.recv().await, None); } +#[tokio::test] +async fn queued_user_message_after_turn_ended_starts_the_next_turn() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push(text_response("first turn")); + provider.push(text_response("queued turn")); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = TransportRegistry::new( + store, + Arc::new(provider.clone()), + PromptCompiler::new(assets), + ); + let handle = registry.get_or_create("queued-after-turn").await.unwrap(); + let mut output = handle.subscribe(); + handle + .command(TransportCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "queued-after-turn", + "queued-after-turn-conversation", + )), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + wait_for_turn_ended(&handle, &mut output, &mut append_seqno).await; + assert_transport_remains_open(&handle, &mut output, &mut append_seqno).await; + + cursor_server::api::cursor::bidi::append( + ®istry, + cursor_server::api::cursor::bidi::DecodedAppend { + request_id: "queued-after-turn".into(), + seqno: append_seqno, + message: runtime_user_message(), + }, + None, + ) + .await + .unwrap(); + append_seqno += 1; + + let mut text = String::new(); + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("queued turn closed without 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 { + text.push_str(&delta.text); + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + + assert!(text.contains("queued turn")); + let requests = provider.requests(); + assert_eq!(requests.len(), 2); + assert_eq!( + &requests[1].history[..requests[0].history.len()], + requests[0].history.as_slice(), + "queued continuation must preserve the first provider request as a prefix" + ); + let history = serde_json::to_string(&requests[1].history).unwrap(); + assert!(history.contains("queued follow-up")); +} + #[tokio::test] async fn runtime_user_message_action_interrupts_and_continues_with_new_message() { let (_directory, store) = fixtures::temp_store().await; @@ -1732,6 +1811,55 @@ async fn run_to_end( } } +async fn wait_for_turn_ended( + handle: &cursor_server::cursor::TransportHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, +) { + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("RunSSE closed before turnEnded"); + 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(); + let turn_ended = matches!( + server.message, + Some(pb::agent_server_message::Message::InteractionUpdate( + pb::InteractionUpdate { + message: Some(pb::interaction_update::Message::TurnEnded(_)), + } + )) + ); + acknowledge_kv(handle, append_seqno, &frame).await; + if turn_ended { + return; + } + } +} + +async fn assert_transport_remains_open( + handle: &cursor_server::cursor::TransportHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, +) { + let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(100); + loop { + let remaining = deadline.saturating_duration_since(tokio::time::Instant::now()); + let Ok(Some(frame)) = tokio::time::timeout(remaining, output.recv()).await else { + return; + }; + let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + assert_eq!( + flags & connect::END_STREAM_FLAG, + 0, + "turnEnded closed the transport before the queued action arrived" + ); + acknowledge_kv(handle, append_seqno, &frame).await; + } +} + async fn acknowledge_kv( handle: &cursor_server::cursor::TransportHandle, append_seqno: &mut i64,