mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
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:
@@ -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)?;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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(®istry, &request.request_id).await
|
run_sse::stream(®istry, &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(®istry, 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(®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());
|
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)
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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() {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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,
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)))
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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, ¤t, 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, ¤t, 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;
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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;
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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,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,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;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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();
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user