feat: enhance bidi request handling and observability tracing

- Updated the `append` function to include a flag for replacing closing requests, improving request handling.
- Refactored the `run_sse_handler` and `bidi_handler` functions to utilize a new tracing mechanism, enhancing observability.
- Introduced a new `trace_outcome` function to standardize tracing outcomes for requests.
- Removed the `CursorTraceRecorder` in favor of a new `CursorTraceService` for better performance and non-blocking behavior.
- Added tests to validate the new tracing functionality and ensure correct behavior during request processing.
This commit is contained in:
leookun
2026-09-02 01:04:54 +08:00
parent 2dad593263
commit 5cdf642dd1
25 changed files with 1522 additions and 458 deletions
+5 -1
View File
@@ -184,7 +184,11 @@ pub async fn append(
request: DecodedAppend, request: DecodedAppend,
parent: Option<TransportParent>, parent: Option<TransportParent>,
) -> Result<ai::BidiAppendResponse> { ) -> Result<ai::BidiAppendResponse> {
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() { if let Some(conversation_id) = request.conversation_id() {
handle.set_conversation_id(conversation_id)?; handle.set_conversation_id(conversation_id)?;
} }
+84 -26
View File
@@ -19,9 +19,7 @@ use crate::{
connect, connect,
proto::{agent::v1 as agent, aiserver::v1 as ai}, proto::{agent::v1 as agent, aiserver::v1 as ai},
}, },
services::{ services::{account, analytics, knowledge, model_catalog, tab},
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
},
transport::{TransportParent, TransportRegistry}, transport::{TransportParent, TransportRegistry},
}, },
Result, Result,
@@ -123,16 +121,13 @@ async fn run_sse_handler(
let (parts, body) = buffered(request).await?; let (parts, body) = buffered(request).await?;
let request: agent::BidiRequestId = connect::decode_unary(&body)?; let request: agent::BidiRequestId = connect::decode_unary(&body)?;
let route = registry.wait_route(&request.request_id).await; let route = registry.wait_route(&request.request_id).await;
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await; let trace = registry.trace(&request.request_id);
if let Some(trace) = &trace { trace.resume();
trace trace.request(
.request( "run_sse_request",
"run_sse_request", body.clone(),
&body, serde_json::json!({"request_id": request.request_id}),
serde_json::json!({"request_id": request.request_id}), );
)
.await;
}
match route { match route {
crate::cursor::transport::TransportRoute::Local => { crate::cursor::transport::TransportRoute::Local => {
run_sse::stream(&registry, &request.request_id).await run_sse::stream(&registry, &request.request_id).await
@@ -143,7 +138,14 @@ async fn run_sse_handler(
Request::from_parts(parts, Body::from(body)), Request::from_parts(parts, Body::from(body)),
) )
.await?; .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 first_model = decoded.model_id().map(str::to_owned);
let conversation_id = decoded.conversation_id().map(str::to_owned); let conversation_id = decoded.conversation_id().map(str::to_owned);
let trace_metadata = decoded.trace_metadata(); let trace_metadata = decoded.trace_metadata();
let trace = registry.trace(&decoded.request_id);
let local = if let Some(model_id) = decoded.model_id() { let local = if let Some(model_id) = decoded.model_id() {
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。 // 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX) 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 { } else if registry.upstream(&decoded.request_id).await {
false false
} else { } else {
trace.resume();
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, false, "missing_transport", None),
);
return Err(crate::Error::Protocol( return Err(crate::Error::Protocol(
"first BidiAppend message must select a model".into(), "first BidiAppend message must select a model".into(),
)); ));
}; };
let trace = if first_model.is_some() { if first_model.is_some() {
CursorTraceRecorder::begin( trace.begin(
registry.store().clone(),
&decoded.request_id,
conversation_id.as_deref(), conversation_id.as_deref(),
if local { if local {
"local_byok" "local_byok"
@@ -198,26 +205,61 @@ async fn bidi_handler(
"cursor_official" "cursor_official"
}, },
first_model.as_deref(), first_model.as_deref(),
) );
.await
} else { } else {
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await trace.resume();
};
if let Some(trace) = &trace {
trace.request("bidi_request", &body, trace_metadata).await;
} }
if !local { if !local {
if first_model.is_some() { if first_model.is_some() {
registry.mark_upstream(&decoded.request_id).await; registry.mark_upstream(&decoded.request_id).await;
} }
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, true, "upstream", None),
);
return proxy::forward( return proxy::forward(
Extension(proxy), Extension(proxy),
Request::from_parts(parts, Body::from(body)), Request::from_parts(parts, Body::from(body)),
) )
.await; .await;
} }
let parent = parent_headers(&parts.headers)?; let parent = match parent_headers(&parts.headers) {
bidi::append(&registry, decoded, parent).await?; 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(&registry, 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()); let mut response = Response::new(axum::body::Body::empty());
*response.status_mut() = StatusCode::OK; *response.status_mut() = StatusCode::OK;
response.headers_mut().insert( response.headers_mut().insert(
@@ -227,6 +269,22 @@ async fn bidi_handler(
Ok(response) Ok(response)
} }
fn trace_outcome(
mut metadata: serde_json::Value,
accepted: bool,
route_outcome: &str,
error: Option<String>,
) -> 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<Body>) -> Result<(axum::http::request::Parts, Bytes)> { async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
let (parts, body) = request.into_parts(); let (parts, body) = request.into_parts();
let body = to_bytes(body, usize::MAX) let body = to_bytes(body, usize::MAX)
+5 -5
View File
@@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Re
let receiver = handle.subscribe(); let receiver = handle.subscribe();
let trace = handle.trace().cloned(); let trace = handle.trace().cloned();
if let Some(trace) = &trace { if let Some(trace) = &trace {
trace.response_started(StatusCode::OK.as_u16()).await; trace.response_started(StatusCode::OK.as_u16());
} }
let body_stream = local_body_stream(receiver, handle, trace); let body_stream = local_body_stream(receiver, handle, trace);
let mut response = Response::new(Body::from_stream(body_stream)); let mut response = Response::new(Body::from_stream(body_stream));
@@ -133,7 +133,7 @@ pub async fn upstream(
) -> Response<Body> { ) -> Response<Body> {
let (parts, body) = response.into_parts(); let (parts, body) = response.into_parts();
if let Some(trace) = &trace { 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 stream = async_stream::stream! {
let _guard = UpstreamRunGuard { let _guard = UpstreamRunGuard {
@@ -180,15 +180,15 @@ impl TraceStreamSink {
while let Some(event) = receiver.recv().await { while let Some(event) = receiver.recv().await {
match event { match event {
TraceStreamEvent::Chunk(chunk) => { TraceStreamEvent::Chunk(chunk) => {
trace.response_chunk(source, &chunk).await; trace.response_chunk(source, chunk);
} }
TraceStreamEvent::Finish(error) => { TraceStreamEvent::Finish(error) => {
trace.finish(error.as_deref()).await; trace.finish(error.as_deref());
return; return;
} }
} }
} }
trace.finish(None).await; trace.finish(None);
}); });
Self { Self {
sender: Some(sender), sender: Some(sender),
+11 -13
View File
@@ -278,19 +278,17 @@ impl CheckpointBuilder {
), ),
}); });
if let Some(trace) = handle.trace() { if let Some(trace) = handle.trace() {
trace trace.artifact(
.artifact( "checkpoint",
"checkpoint", "byok_server",
"byok_server", &checkpoint.encode_to_vec(),
&checkpoint.encode_to_vec(), serde_json::json!({
serde_json::json!({ "root_message_count": checkpoint.root_prompt_messages_json.len(),
"root_message_count": checkpoint.root_prompt_messages_json.len(), "turn_count": checkpoint.turns.len(),
"turn_count": checkpoint.turns.len(), "pending_tool_call_count": checkpoint.pending_tool_calls.len(),
"pending_tool_call_count": checkpoint.pending_tool_calls.len(), "emit_status": if result.is_ok() { "sent" } else { "error" },
"emit_status": if result.is_ok() { "sent" } else { "error" }, }),
}), );
)
.await;
} }
result result
} }
+2 -14
View File
@@ -12,7 +12,7 @@ use crate::{
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer, protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
tools::runtime::McpRoute, tools::runtime::McpRoute,
}, },
model::ToolDefinition, model::{normalize_tool_name, ToolDefinition},
store::BlobId, store::BlobId,
Error, Result, Error, Result,
}; };
@@ -493,7 +493,7 @@ pub fn dynamic_mcp(
})?), })?),
}; };
let parameters = normalize_mcp_parameters(&wire.name, parameters)?; 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 { let definition = ToolDefinition {
name: name.clone(), name: name.clone(),
description: wire.description.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 { fn prost_value(value: &prost_types::Value) -> Value {
use prost_types::value::Kind; use prost_types::value::Kind;
match value.kind.as_ref() { match value.kind.as_ref() {
+1 -3
View File
@@ -117,9 +117,7 @@ pub(crate) async fn prepare(
"selected_source": "root_prompt_messages_json", "selected_source": "root_prompt_messages_json",
}); });
let encoded = serde_json::to_vec(&summary)?; let encoded = serde_json::to_vec(&summary)?;
trace trace.artifact("history_projection", "byok_server", &encoded, summary);
.artifact("history_projection", "byok_server", &encoded, summary)
.await;
} }
let mut request_context = context::hydrate(request, context_sync).await?; let mut request_context = context::hydrate(request, context_sync).await?;
if let Some(rules_dir) = local_rules_dir { if let Some(rules_dir) = local_rules_dir {
+12 -2
View File
@@ -1,6 +1,13 @@
//! Defines commands accepted by a Conversation runtime. //! 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 TransportFinish {
Success,
Failed(Error),
Cancelled,
}
#[derive(Debug)] #[derive(Debug)]
pub enum TransportCommand { pub enum TransportCommand {
@@ -8,6 +15,9 @@ pub enum TransportCommand {
seqno: i64, seqno: i64,
message: Box<pb::AgentClientMessage>, message: Box<pb::AgentClientMessage>,
}, },
RunFinished {
generation: u64,
finish: TransportFinish,
},
Disconnect, Disconnect,
Close,
} }
+10 -12
View File
@@ -34,7 +34,7 @@ use crate::{
Error, Result, Error, Result,
}; };
use super::{CompiledMessages, ConversationRegistry, MessageDelivery}; use super::{CompiledMessages, ConversationRegistry, MessageDelivery, TransportFinish};
use crate::cursor::transport::TransportHandle; use crate::cursor::transport::TransportHandle;
pub struct ConversationOutput { pub struct ConversationOutput {
@@ -110,7 +110,7 @@ impl ConversationOutput {
} }
} }
pub async fn run(mut self) -> Result<()> { pub async fn run(mut self) -> Result<TransportFinish> {
let result = self.run_inner().await; let result = self.run_inner().await;
if let Err(error) = &result { if let Err(error) = &result {
if !self.superseded.is_cancelled() { if !self.superseded.is_cancelled() {
@@ -143,7 +143,7 @@ impl ConversationOutput {
result result
} }
async fn run_inner(&mut self) -> Result<()> { async fn run_inner(&mut self) -> Result<TransportFinish> {
if self.context.compacting { if self.context.compacting {
self.handle.emit(&events::summary_started())?; self.handle.emit(&events::summary_started())?;
} }
@@ -175,7 +175,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() { if self.superseded.is_cancelled() {
worker.abort(); worker.abort();
self.abort_execs().await; self.abort_execs().await;
return Ok(()); return Ok(TransportFinish::Cancelled);
} }
let input = if let Ok(action) = self.runtime_actions.try_recv() { let input = if let Ok(action) = self.runtime_actions.try_recv() {
Input::RuntimeAction(Some(Box::new(action))) Input::RuntimeAction(Some(Box::new(action)))
@@ -187,7 +187,7 @@ impl ConversationOutput {
_ = self.superseded.cancelled() => { _ = self.superseded.cancelled() => {
worker.abort(); worker.abort();
self.abort_execs().await; self.abort_execs().await;
return Ok(()); return Ok(TransportFinish::Cancelled);
} }
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)), action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
event = self.core.events.recv() => Input::Event(event), event = self.core.events.recv() => Input::Event(event),
@@ -719,7 +719,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() { if self.superseded.is_cancelled() {
worker.abort(); worker.abort();
self.abort_execs().await; self.abort_execs().await;
return Ok(()); return Ok(TransportFinish::Cancelled);
} }
return match outcome { return match outcome {
RunOutcome::Completed => { RunOutcome::Completed => {
@@ -735,8 +735,7 @@ impl ConversationOutput {
for _ in 0..3 { for _ in 0..3 {
self.checkpoint.publish(&self.handle, &checkpoint).await?; self.checkpoint.publish(&self.handle, &checkpoint).await?;
} }
finish_success(&self.handle); return Ok(TransportFinish::Success);
return Ok(());
} }
let checkpoints = final_checkpoint.take().ok_or_else(|| { let checkpoints = final_checkpoint.take().ok_or_else(|| {
Error::Protocol("Completed without final state".into()) Error::Protocol("Completed without final state".into())
@@ -752,18 +751,17 @@ impl ConversationOutput {
ttft_breakdown: None, ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)), message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
})?; })?;
finish_success(&self.handle); Ok(TransportFinish::Success)
Ok(())
} }
RunOutcome::Cancelled => { RunOutcome::Cancelled => {
worker.abort(); worker.abort();
self.abort_execs().await; self.abort_execs().await;
finish_cancelled(&self.handle) Ok(TransportFinish::Cancelled)
} }
RunOutcome::Failed(failure) => { RunOutcome::Failed(failure) => {
worker.abort(); worker.abort();
self.abort_execs().await; self.abort_execs().await;
finish_failed(&self.handle, &cursor_error(failure)) Ok(TransportFinish::Failed(cursor_error(failure)))
} }
}; };
} }
+132 -31
View File
@@ -23,13 +23,14 @@ use crate::{
use super::{ use super::{
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies, CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
ConversationRegistry, MessageDelivery, TransportCommand, ConversationRegistry, MessageDelivery, TransportCommand, TransportFinish,
}; };
pub struct ConversationRuntime; pub struct ConversationRuntime;
#[derive(Clone)] #[derive(Clone)]
struct RunGeneration { struct RunGeneration {
id: u64,
superseded: CancellationToken, superseded: CancellationToken,
finished: CancellationToken, finished: CancellationToken,
run: Arc<parking_lot::Mutex<Option<RunHandle>>>, run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
@@ -41,6 +42,14 @@ struct RunGeneration {
struct FinishGeneration(CancellationToken); struct FinishGeneration(CancellationToken);
struct TransportActorGuard(TransportHandle);
impl Drop for TransportActorGuard {
fn drop(&mut self) {
self.0.close_transport();
}
}
impl Drop for FinishGeneration { impl Drop for FinishGeneration {
fn drop(&mut self) { fn drop(&mut self) {
self.0.cancel(); self.0.cancel();
@@ -54,6 +63,7 @@ impl ConversationRuntime {
mut receiver: mpsc::Receiver<TransportCommand>, mut receiver: mpsc::Receiver<TransportCommand>,
) { ) {
tokio::spawn(async move { tokio::spawn(async move {
let _actor_guard = TransportActorGuard(handle.clone());
let dependencies = registry.dependencies().clone(); let dependencies = registry.dependencies().clone();
let blob_sync = BlobSynchronizer::new( let blob_sync = BlobSynchronizer::new(
handle.request_id().into(), handle.request_id().into(),
@@ -65,19 +75,47 @@ impl ConversationRuntime {
let context_sync = let context_sync =
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone()); RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
let mut current = None::<RunGeneration>; let mut current = None::<RunGeneration>;
let mut next_generation = 1_u64;
let mut pending_finish = None::<(u64, TransportFinish)>;
let mut draining = false;
loop { loop {
let command = match receiver.recv().await { let command = if draining {
Some(command) => command, if !handle.admissions_drained() {
None => { tokio::select! {
handle.mark_disconnected(); command = receiver.recv() => match command {
if let Some(generation) = current.as_ref() { Some(command) => command,
generation.superseded.cancel(); None => {
if let Some(run) = generation.run.lock().clone() { finish_pending(&handle, &current, pending_finish.take());
run.cancel(); 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, &current, pending_finish.take());
break;
} }
} }
super::finish_cancelled(&handle).ok(); }
break; } 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 { match command {
@@ -95,8 +133,16 @@ impl ConversationRuntime {
super::finish_cancelled(&handle).ok(); super::finish_cancelled(&handle).ok();
break; break;
} }
TransportCommand::Close => { TransportCommand::RunFinished { generation, finish } => {
break; if !current
.as_ref()
.is_some_and(|current| current.id == generation)
{
continue;
}
handle.begin_close();
pending_finish = Some((generation, finish));
draining = true;
} }
TransportCommand::Append { seqno, message } => { TransportCommand::Append { seqno, message } => {
for (_seqno, message) in inbox.push(seqno, *message) { for (_seqno, message) in inbox.push(seqno, *message) {
@@ -105,6 +151,11 @@ impl ConversationRuntime {
Some(pb::agent_client_message::Message::RunRequest( Some(pb::agent_client_message::Message::RunRequest(
request, request,
)) => { )) => {
if draining {
handle.reopen();
draining = false;
pending_finish = None;
}
if let Some(conversation_id) = if let Some(conversation_id) =
request.conversation_id.as_deref() request.conversation_id.as_deref()
{ {
@@ -117,8 +168,6 @@ impl ConversationRuntime {
"invalid Cursor conversation id" "invalid Cursor conversation id"
); );
let _ = super::finish_failed(&handle, &error); let _ = super::finish_failed(&handle, &error);
let _ =
handle.command(TransportCommand::Close).await;
return; return;
} }
} }
@@ -150,6 +199,7 @@ impl ConversationRuntime {
dependencies.web_cache.clone(), dependencies.web_cache.clone(),
); );
let generation = RunGeneration { let generation = RunGeneration {
id: next_generation,
superseded: CancellationToken::new(), superseded: CancellationToken::new(),
finished: CancellationToken::new(), finished: CancellationToken::new(),
run: Arc::new(parking_lot::Mutex::new(None)), run: Arc::new(parking_lot::Mutex::new(None)),
@@ -158,6 +208,7 @@ impl ConversationRuntime {
tool_runtime, tool_runtime,
tools, tools,
}; };
next_generation = next_generation.saturating_add(1);
current = Some(generation.clone()); current = Some(generation.clone());
spawn_run_request( spawn_run_request(
registry.clone(), registry.clone(),
@@ -428,6 +479,34 @@ impl ConversationRuntime {
} }
} }
fn finish_pending(
handle: &TransportHandle,
current: &Option<RunGeneration>,
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)] #[allow(clippy::too_many_arguments)]
fn spawn_run_request( fn spawn_run_request(
registry: ConversationRegistry, registry: ConversationRegistry,
@@ -487,8 +566,12 @@ fn spawn_run_request(
%error, %error,
"failed to prepare Cursor Run" "failed to prepare Cursor Run"
); );
let _ = super::finish_failed(&handle, &error); let _ = handle
let _ = handle.command(TransportCommand::Close).await; .command(TransportCommand::RunFinished {
generation: generation.id,
finish: TransportFinish::Failed(error),
})
.await;
return; return;
} }
}; };
@@ -520,8 +603,12 @@ fn spawn_run_request(
{ {
CommandResult::Applied | CommandResult::Duplicate => { CommandResult::Applied | CommandResult::Duplicate => {
if !generation.superseded.is_cancelled() { if !generation.superseded.is_cancelled() {
super::finish_success(&handle); let _ = handle
let _ = handle.command(TransportCommand::Close).await; .command(TransportCommand::RunFinished {
generation: generation.id,
finish: TransportFinish::Success,
})
.await;
} }
return; return;
} }
@@ -545,8 +632,12 @@ fn spawn_run_request(
} }
CommandResult::StaleTarget => { CommandResult::StaleTarget => {
if !generation.superseded.is_cancelled() { if !generation.superseded.is_cancelled() {
super::finish_success(&handle); let _ = handle
let _ = handle.command(TransportCommand::Close).await; .command(TransportCommand::RunFinished {
generation: generation.id,
finish: TransportFinish::Success,
})
.await;
} }
return; return;
} }
@@ -605,16 +696,21 @@ fn spawn_run_request(
tool_runtime: generation.tool_runtime.clone(), tool_runtime: generation.tool_runtime.clone(),
}, },
); );
if let Err(error) = output.run().await { let finish = match output.run().await {
if !generation.superseded.is_cancelled() { Ok(finish) => finish,
tracing::error!( Err(error) => {
request_id = handle.request_id(), if generation.superseded.is_cancelled() {
%error, TransportFinish::Cancelled
"Cursor session failed" } else {
); tracing::error!(
let _ = super::finish_failed(&handle, &error); request_id = handle.request_id(),
%error,
"Cursor session failed"
);
TransportFinish::Failed(error)
}
} }
} };
let _ = core_run.await; let _ = core_run.await;
registry.release(&conversation_id, &run_id).await; registry.release(&conversation_id, &run_id).await;
if generation if generation
@@ -626,7 +722,12 @@ fn spawn_run_request(
*generation.run.lock() = None; *generation.run.lock() = None;
} }
if !generation.superseded.is_cancelled() { if !generation.superseded.is_cancelled() {
let _ = handle.command(TransportCommand::Close).await; let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish,
})
.await;
} }
}); });
} }
+53 -63
View File
@@ -76,22 +76,20 @@ impl BlobSynchronizer {
let id = self.inner.store.put_blob(data, edges).await?; let id = self.inner.store.put_blob(data, edges).await?;
let result = self.ensure_set(&id, data).await; let result = self.ensure_set(&id, data).await;
if let Some(trace) = self.inner.handle.trace() { if let Some(trace) = self.inner.handle.trace() {
trace trace.linked_blob(
.linked_blob( "blob_set",
"blob_set", "byok_server",
"byok_server", &id,
&id, serde_json::json!({
serde_json::json!({ "byte_count": data.len(),
"byte_count": data.len(), "status": if result.is_ok() { "acknowledged" } else { "error" },
"status": if result.is_ok() { "acknowledged" } else { "error" }, "error": result.as_ref().err().map(ToString::to_string),
"error": result.as_ref().err().map(ToString::to_string), "edges": edges.iter().map(|edge| serde_json::json!({
"edges": edges.iter().map(|edge| serde_json::json!({ "child_blob_id": edge.child.to_base64(),
"child_blob_id": edge.child.to_base64(), "field_name": edge.field_name,
"field_name": edge.field_name, })).collect::<Vec<_>>(),
})).collect::<Vec<_>>(), }),
}), );
)
.await;
} }
result?; result?;
Ok(id) Ok(id)
@@ -144,18 +142,16 @@ impl BlobSynchronizer {
pub async fn get(&self, blob_id: &BlobId) -> Result<Option<Vec<u8>>> { pub async fn get(&self, blob_id: &BlobId) -> Result<Option<Vec<u8>>> {
if let Some(data) = self.inner.store.get_blob(blob_id).await? { if let Some(data) = self.inner.store.get_blob(blob_id).await? {
if let Some(trace) = self.inner.handle.trace() { if let Some(trace) = self.inner.handle.trace() {
trace trace.linked_blob(
.linked_blob( "blob_get",
"blob_get", "byok_server",
"byok_server", blob_id,
blob_id, serde_json::json!({
serde_json::json!({ "byte_count": data.len(),
"byte_count": data.len(), "source": "local_store",
"source": "local_store", "status": "found",
"status": "found", }),
}), );
)
.await;
} }
return Ok(Some(data)); return Ok(Some(data));
} }
@@ -194,45 +190,39 @@ impl BlobSynchronizer {
if let Some(trace) = self.inner.handle.trace() { if let Some(trace) = self.inner.handle.trace() {
match &result { match &result {
Ok(Some(data)) => { Ok(Some(data)) => {
trace trace.linked_blob(
.linked_blob( "blob_get",
"blob_get", "cursor_client",
"cursor_client", blob_id,
blob_id, serde_json::json!({
serde_json::json!({ "byte_count": data.len(),
"byte_count": data.len(), "source": "cursor_client",
"source": "cursor_client", "status": "found",
"status": "found", }),
}), );
)
.await;
} }
Ok(None) => { Ok(None) => {
trace trace.artifact(
.artifact( "blob_get",
"blob_get", "cursor_client",
"cursor_client", &[],
&[], serde_json::json!({
serde_json::json!({ "blob_id": blob_id.to_base64(),
"blob_id": blob_id.to_base64(), "status": "missing",
"status": "missing", }),
}), );
)
.await;
} }
Err(error) => { Err(error) => {
trace trace.artifact(
.artifact( "blob_get",
"blob_get", "cursor_client",
"cursor_client", &[],
&[], serde_json::json!({
serde_json::json!({ "blob_id": blob_id.to_base64(),
"blob_id": blob_id.to_base64(), "status": "error",
"status": "error", "error": error.to_string(),
"error": error.to_string(), }),
}), );
)
.await;
} }
} }
} }
-225
View File
@@ -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<Mutex<TraceChunkBuffer>>,
finished: Arc<AtomicBool>,
}
#[derive(Default)]
struct TraceChunkBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
first_chunk_at: Option<Instant>,
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<Self> {
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<Self> {
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(())
}
}
@@ -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<AtomicU8>,
conversation_id: Option<String>,
route: String,
model_id: Option<String>,
},
Resume {
request_id: String,
activation: Arc<AtomicU8>,
},
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<String>,
},
}
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,
}
}
}
@@ -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<TraceEvent>,
}
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<str>,
sender: mpsc::Sender<TraceEvent>,
finished: Arc<AtomicBool>,
activation: Arc<AtomicU8>,
}
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"
);
}
}
}
@@ -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<BufferedCursorTraceChunk>,
bytes: usize,
}
struct BufferedRequest {
artifact_type: String,
data: bytes::Bytes,
metadata: serde_json::Value,
}
#[derive(Default)]
struct RequestOrder {
next: i64,
pending: BTreeMap<i64, Vec<BufferedRequest>>,
}
pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver<TraceEvent>) {
let mut states = HashMap::<String, TraceState>::new();
let mut buffers = HashMap::<String, ResponseBuffer>::new();
let mut request_orders = HashMap::<String, RequestOrder>::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<String, TraceState>,
buffers: &mut HashMap<String, ResponseBuffer>,
request_orders: &mut HashMap<String, RequestOrder>,
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<String, RequestOrder>,
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<String, RequestOrder>,
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<String, TraceState>,
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<String, ResponseBuffer>, 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<String, ResponseBuffer>) {
let request_ids = buffers.keys().cloned().collect::<Vec<_>>();
for request_id in request_ids {
flush_one(store, buffers, &request_id).await;
}
}
+43 -7
View File
@@ -15,7 +15,7 @@ use crate::{
Error, Result, Error, Result,
}; };
use super::OutputHub; use super::{OutputHub, TransportAdmission, TransportLifecycle};
#[derive(Clone, Debug, PartialEq, Eq)] #[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransportParent { pub struct TransportParent {
@@ -30,7 +30,8 @@ pub struct TransportHandle {
output: Arc<OutputHub>, output: Arc<OutputHub>,
conversation_id: Arc<OnceLock<String>>, conversation_id: Arc<OnceLock<String>>,
parent: Arc<OnceLock<TransportParent>>, parent: Arc<OnceLock<TransportParent>>,
trace: Option<CursorTraceRecorder>, trace: CursorTraceRecorder,
lifecycle: TransportLifecycle,
disconnect: CancellationToken, disconnect: CancellationToken,
} }
@@ -39,7 +40,7 @@ impl TransportHandle {
request_id: String, request_id: String,
commands: mpsc::Sender<TransportCommand>, commands: mpsc::Sender<TransportCommand>,
output: Arc<OutputHub>, output: Arc<OutputHub>,
trace: Option<CursorTraceRecorder>, trace: CursorTraceRecorder,
) -> Self { ) -> Self {
Self { Self {
request_id, request_id,
@@ -48,6 +49,7 @@ impl TransportHandle {
conversation_id: Arc::new(OnceLock::new()), conversation_id: Arc::new(OnceLock::new()),
parent: Arc::new(OnceLock::new()), parent: Arc::new(OnceLock::new()),
trace, trace,
lifecycle: TransportLifecycle::new(),
disconnect: CancellationToken::new(), disconnect: CancellationToken::new(),
} }
} }
@@ -127,12 +129,46 @@ impl TransportHandle {
self.output.close() self.output.close()
} }
pub(crate) async fn wait_closed(&self) { pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
self.output.wait_closed().await; Some(&self.trace)
} }
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> { pub(crate) fn accepting_appends(&self) -> bool {
self.trace.as_ref() self.lifecycle.is_open()
}
pub(crate) fn admit(&self) -> Result<TransportAdmission> {
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 { pub(crate) fn disconnect_token(&self) -> CancellationToken {
+170
View File
@@ -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<LifecycleInner>,
}
struct LifecycleInner {
state: parking_lot::Mutex<LifecycleState>,
admissions_drained: Notify,
closed: Notify,
}
struct LifecycleState {
state: TransportState,
admissions: usize,
}
pub(crate) struct TransportAdmission {
inner: Arc<LifecycleInner>,
}
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<TransportAdmission> {
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());
}
}
+2
View File
@@ -2,10 +2,12 @@
mod handle; mod handle;
mod inbox; mod inbox;
mod lifecycle;
mod output; mod output;
mod registry; mod registry;
pub use handle::*; pub use handle::*;
pub use inbox::*; pub use inbox::*;
pub(crate) use lifecycle::*;
pub use output::*; pub use output::*;
pub use registry::*; pub use registry::*;
+1 -13
View File
@@ -1,12 +1,11 @@
//! Buffers, replays, broadcasts, and atomically closes downstream output. //! Buffers, replays, broadcasts, and atomically closes downstream output.
use bytes::Bytes; use bytes::Bytes;
use tokio::sync::{mpsc, Notify}; use tokio::sync::mpsc;
#[derive(Default)] #[derive(Default)]
pub struct OutputHub { pub struct OutputHub {
state: parking_lot::Mutex<OutputState>, state: parking_lot::Mutex<OutputState>,
closed: Notify,
} }
#[derive(Default)] #[derive(Default)]
@@ -49,17 +48,6 @@ impl OutputHub {
state.closed = true; state.closed = true;
state.subscribers.clear(); state.subscribers.clear();
drop(state); drop(state);
self.closed.notify_waiters();
true true
} }
pub async fn wait_closed(&self) {
loop {
let notified = self.closed.notified();
if self.state.lock().closed {
return;
}
notified.await;
}
}
} }
+76 -19
View File
@@ -1,13 +1,19 @@
//! Maps request IDs to active transport handles. //! 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 tokio::sync::{mpsc, Mutex, Notify};
use crate::{ use crate::{
cursor::{ cursor::{
conversation::ConversationRegistry, prompting::PromptCompiler, conversation::ConversationRegistry, prompting::PromptCompiler,
services::observability::CursorTraceRecorder, services::observability::CursorTraceService,
}, },
plugin::PluginRegistry, plugin::PluginRegistry,
provider::Provider, provider::Provider,
@@ -24,15 +30,23 @@ pub struct TransportRegistry {
} }
struct RegistryInner { struct RegistryInner {
local: Mutex<HashMap<String, TransportHandle>>, local: Mutex<HashMap<String, LocalTransport>>,
next_local_generation: AtomicU64,
upstream: Mutex<HashMap<String, u64>>, upstream: Mutex<HashMap<String, u64>>,
route_changed: Notify, route_changed: Notify,
store: Store, store: Store,
traces: CursorTraceService,
web_cache: WebCache, web_cache: WebCache,
plugins: Option<PluginRegistry>, plugins: Option<PluginRegistry>,
conversations: ConversationRegistry, conversations: ConversationRegistry,
} }
#[derive(Clone)]
struct LocalTransport {
generation: u64,
handle: TransportHandle,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)] #[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TransportRoute { pub enum TransportRoute {
Local, Local,
@@ -99,8 +113,10 @@ impl TransportRegistry {
Self { Self {
inner: Arc::new(RegistryInner { inner: Arc::new(RegistryInner {
local: Mutex::new(HashMap::new()), local: Mutex::new(HashMap::new()),
next_local_generation: AtomicU64::new(1),
upstream: Mutex::new(HashMap::new()), upstream: Mutex::new(HashMap::new()),
route_changed: Notify::new(), route_changed: Notify::new(),
traces: CursorTraceService::new(store.clone()),
conversations: ConversationRegistry::new( conversations: ConversationRegistry::new(
store.clone(), store.clone(),
provider, provider,
@@ -119,6 +135,13 @@ impl TransportRegistry {
&self.inner.store &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 { pub fn web_cache(&self) -> &WebCache {
&self.inner.web_cache &self.inner.web_cache
} }
@@ -132,18 +155,37 @@ impl TransportRegistry {
} }
pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> { pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> {
if let Some(handle) = self.inner.local.lock().await.get(request_id).cloned() { self.get_or_create_for_append(request_id, false).await
return Ok(handle); }
pub(crate) async fn get_or_create_for_append(
&self,
request_id: &str,
replace_closing: bool,
) -> Result<TransportHandle> {
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 (commands, receiver) = mpsc::channel(128);
let output = Arc::new(OutputHub::default()); let output = Arc::new(OutputHub::default());
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await; let trace = self.inner.traces.recorder(request_id);
let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace); trace.resume();
let mut local = self.inner.local.lock().await; let handle = TransportHandle::new(request_id.into(), commands, output, trace);
if let Some(existing) = local.get(request_id).cloned() { let generation = self
return Ok(existing); .inner
} .next_local_generation
local.insert(request_id.into(), handle.clone()); .fetch_add(1, Ordering::Relaxed);
local.insert(
request_id.into(),
LocalTransport {
generation,
handle: handle.clone(),
},
);
drop(local); drop(local);
self.inner.route_changed.notify_waiters(); self.inner.route_changed.notify_waiters();
self.inner self.inner
@@ -152,17 +194,29 @@ impl TransportRegistry {
let registry = Arc::downgrade(&self.inner); let registry = Arc::downgrade(&self.inner);
let request_id = request_id.to_string(); let request_id = request_id.to_string();
let lifecycle = handle.clone();
tokio::spawn(async move { tokio::spawn(async move {
output.wait_closed().await; lifecycle.wait_transport_closed().await;
if let Some(registry) = registry.upgrade() { 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) Ok(handle)
} }
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> { pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
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) { pub async fn mark_upstream(&self, request_id: &str) {
@@ -206,10 +260,13 @@ impl TransportRegistry {
self.inner.conversations.shutdown().await; self.inner.conversations.shutdown().await;
let handles = std::mem::take(&mut *self.inner.local.lock().await); let handles = std::mem::take(&mut *self.inner.local.lock().await);
self.inner.upstream.lock().await.clear(); self.inner.upstream.lock().await.clear();
for handle in handles.into_values() { for transport in handles.into_values() {
handle.disconnect().await; transport.handle.disconnect().await;
let _ = let _ = tokio::time::timeout(
tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await; std::time::Duration::from_secs(2),
transport.handle.wait_transport_closed(),
)
.await;
} }
} }
} }
+55 -6
View File
@@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize};
use crate::{Error, Result}; use crate::{Error, Result};
use super::{ use super::{
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent, normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role,
ToolResultContent, ToolCallContent, ToolResultContent,
}; };
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
@@ -91,7 +91,7 @@ fn project_tool_round(
"tool round repeats provider replay state".into(), "tool round repeats provider replay state".into(),
)); ));
} }
calls.extend(part_calls.iter().cloned()); calls.extend(part_calls.iter().map(normalized_tool_call));
cursor += 1; cursor += 1;
while cursor < messages.len() { while cursor < messages.len() {
@@ -139,7 +139,7 @@ fn project_tool_round(
.map(|(message_id, result)| ProjectedMessage { .map(|(message_id, result)| ProjectedMessage {
message_id, message_id,
role: Role::Tool, role: Role::Tool,
content: ProjectedContent::ToolResult(result), content: ProjectedContent::ToolResult(normalized_tool_result(&result)),
}), }),
); );
Ok(Some((output, cursor))) Ok(Some((output, cursor)))
@@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
text: text.clone(), text: text.clone(),
thinking: thinking.clone(), thinking: thinking.clone(),
replay_state: replay_state.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 { ProjectedMessage {
message_id: message.message_id.clone(), message_id: message.message_id.clone(),
@@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
content, 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");
}
}
+18
View File
@@ -4,6 +4,24 @@ use serde_json::Value;
use super::ProviderReplayState; 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::<String>();
if normalized.is_empty() {
"_".into()
} else {
normalized
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
pub struct ToolDefinition { pub struct ToolDefinition {
pub name: String, pub name: String,
+30 -1
View File
@@ -6,7 +6,7 @@ use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::{ use crate::{
model::{ProviderReplayState, ToolCall, Usage}, model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage},
provider::{FinishReason, ModelEvent, ProviderStream}, provider::{FinishReason, ModelEvent, ProviderStream},
}; };
@@ -165,6 +165,7 @@ pub async fn consume_model_cycle(
call_id, call_id,
name, name,
} => { } => {
let name = normalize_tool_name(&name);
let Some(model_call_id) = model_call_id.as_ref() else { let Some(model_call_id) = model_call_id.as_ref() else {
return Err(failure( return Err(failure(
RunFailure::Protocol("provider emitted content before Start".into()), RunFailure::Protocol("provider emitted content before Start".into()),
@@ -406,6 +407,34 @@ mod tests {
}; };
use tokio_stream::wrappers::ReceiverStream; 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] #[tokio::test]
async fn usage_is_forwarded_before_the_provider_call_finishes() { async fn usage_is_forwarded_before_the_provider_call_finishes() {
let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4); let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4);
+34 -17
View File
@@ -62,6 +62,40 @@ impl Store {
.await?) .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( pub async fn append_cursor_trace_artifact(
&self, &self,
request_id: &str, request_id: &str,
@@ -144,23 +178,6 @@ impl Store {
Ok(()) 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<()> { pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> {
let now = now_ms(); let now = now_ms();
let _write = self.writes.lock().await; let _write = self.writes.lock().await;
+129
View File
@@ -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::<Vec<_>>();
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::<i64>()
);
}
#[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());
}
@@ -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();
}