Merge branch 'main' of github.com:leookun/cursor-byok

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