mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:40:50 +08:00
Merge branch 'main' of github.com:leookun/cursor-byok
This commit is contained in:
@@ -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)?;
|
||||
}
|
||||
|
||||
@@ -19,9 +19,7 @@ use crate::{
|
||||
connect,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
},
|
||||
services::{
|
||||
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
|
||||
},
|
||||
services::{account, analytics, knowledge, model_catalog, tab},
|
||||
transport::{TransportParent, TransportRegistry},
|
||||
},
|
||||
Result,
|
||||
@@ -123,16 +121,13 @@ async fn run_sse_handler(
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
|
||||
let route = registry.wait_route(&request.request_id).await;
|
||||
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
|
||||
if let Some(trace) = &trace {
|
||||
trace
|
||||
.request(
|
||||
"run_sse_request",
|
||||
&body,
|
||||
serde_json::json!({"request_id": request.request_id}),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let trace = registry.trace(&request.request_id);
|
||||
trace.resume();
|
||||
trace.request(
|
||||
"run_sse_request",
|
||||
body.clone(),
|
||||
serde_json::json!({"request_id": request.request_id}),
|
||||
);
|
||||
match route {
|
||||
crate::cursor::transport::TransportRoute::Local => {
|
||||
run_sse::stream(®istry, &request.request_id).await
|
||||
@@ -143,7 +138,14 @@ async fn run_sse_handler(
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await?;
|
||||
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
|
||||
Ok(run_sse::upstream(
|
||||
registry,
|
||||
request.request_id,
|
||||
generation,
|
||||
response,
|
||||
Some(trace),
|
||||
)
|
||||
.await)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -159,6 +161,7 @@ async fn bidi_handler(
|
||||
let first_model = decoded.model_id().map(str::to_owned);
|
||||
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
||||
let trace_metadata = decoded.trace_metadata();
|
||||
let trace = registry.trace(&decoded.request_id);
|
||||
let local = if let Some(model_id) = decoded.model_id() {
|
||||
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
|
||||
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
|
||||
@@ -183,14 +186,18 @@ async fn bidi_handler(
|
||||
} else if registry.upstream(&decoded.request_id).await {
|
||||
false
|
||||
} else {
|
||||
trace.resume();
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body.clone(),
|
||||
trace_outcome(trace_metadata, false, "missing_transport", None),
|
||||
);
|
||||
return Err(crate::Error::Protocol(
|
||||
"first BidiAppend message must select a model".into(),
|
||||
));
|
||||
};
|
||||
let trace = if first_model.is_some() {
|
||||
CursorTraceRecorder::begin(
|
||||
registry.store().clone(),
|
||||
&decoded.request_id,
|
||||
if first_model.is_some() {
|
||||
trace.begin(
|
||||
conversation_id.as_deref(),
|
||||
if local {
|
||||
"local_byok"
|
||||
@@ -198,26 +205,61 @@ async fn bidi_handler(
|
||||
"cursor_official"
|
||||
},
|
||||
first_model.as_deref(),
|
||||
)
|
||||
.await
|
||||
);
|
||||
} else {
|
||||
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
|
||||
};
|
||||
if let Some(trace) = &trace {
|
||||
trace.request("bidi_request", &body, trace_metadata).await;
|
||||
trace.resume();
|
||||
}
|
||||
if !local {
|
||||
if first_model.is_some() {
|
||||
registry.mark_upstream(&decoded.request_id).await;
|
||||
}
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body.clone(),
|
||||
trace_outcome(trace_metadata, true, "upstream", None),
|
||||
);
|
||||
return proxy::forward(
|
||||
Extension(proxy),
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let parent = parent_headers(&parts.headers)?;
|
||||
bidi::append(®istry, decoded, parent).await?;
|
||||
let parent = match parent_headers(&parts.headers) {
|
||||
Ok(parent) => parent,
|
||||
Err(error) => {
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(
|
||||
trace_metadata,
|
||||
false,
|
||||
"invalid_parent",
|
||||
Some(error.to_string()),
|
||||
),
|
||||
);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
match bidi::append(®istry, decoded, parent).await {
|
||||
Ok(_) => trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(trace_metadata, true, "local", None),
|
||||
),
|
||||
Err(error) => {
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(
|
||||
trace_metadata,
|
||||
false,
|
||||
"command_rejected",
|
||||
Some(error.to_string()),
|
||||
),
|
||||
);
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
let mut response = Response::new(axum::body::Body::empty());
|
||||
*response.status_mut() = StatusCode::OK;
|
||||
response.headers_mut().insert(
|
||||
@@ -227,6 +269,22 @@ async fn bidi_handler(
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
fn trace_outcome(
|
||||
mut metadata: serde_json::Value,
|
||||
accepted: bool,
|
||||
route_outcome: &str,
|
||||
error: Option<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)
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
))))
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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, ¤t, pending_finish.take());
|
||||
break;
|
||||
}
|
||||
},
|
||||
_ = handle.wait_admissions_drained() => continue,
|
||||
}
|
||||
} else {
|
||||
handle.mark_draining();
|
||||
match receiver.try_recv() {
|
||||
Ok(command) => command,
|
||||
Err(mpsc::error::TryRecvError::Empty)
|
||||
| Err(mpsc::error::TryRecvError::Disconnected) => {
|
||||
finish_pending(&handle, ¤t, pending_finish.take());
|
||||
break;
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
break;
|
||||
}
|
||||
} else if waiting_for_action {
|
||||
tokio::select! {
|
||||
command = receiver.recv() => match command {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
}
|
||||
},
|
||||
_ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
};
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation.id, TransportFinish::Success));
|
||||
draining = true;
|
||||
waiting_for_action = false;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
match receiver.recv().await {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
if let Some(generation) = current.as_ref() {
|
||||
generation.superseded.cancel();
|
||||
if let Some(run) = generation.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
break;
|
||||
}
|
||||
}
|
||||
};
|
||||
match command {
|
||||
@@ -92,11 +156,41 @@ impl ConversationRuntime {
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
let turn_completed = waiting_for_action
|
||||
|| current.as_ref().is_some_and(|generation| {
|
||||
generation
|
||||
.run
|
||||
.lock()
|
||||
.as_ref()
|
||||
.is_none_or(|run| run.phase() != RunPhase::Running)
|
||||
});
|
||||
if turn_completed {
|
||||
super::finish_success(&handle);
|
||||
} else {
|
||||
super::finish_cancelled(&handle).ok();
|
||||
}
|
||||
break;
|
||||
}
|
||||
TransportCommand::Close => {
|
||||
break;
|
||||
TransportCommand::RunFinished { generation, finish } => {
|
||||
if !current
|
||||
.as_ref()
|
||||
.is_some_and(|current| current.id == generation)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
match finish {
|
||||
RunFinish::TurnCompleted => {
|
||||
pending_finish = None;
|
||||
draining = false;
|
||||
waiting_for_action = true;
|
||||
}
|
||||
RunFinish::Transport(finish) => {
|
||||
waiting_for_action = false;
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation, finish));
|
||||
draining = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
TransportCommand::Append { seqno, message } => {
|
||||
for (_seqno, message) in inbox.push(seqno, *message) {
|
||||
@@ -105,6 +199,12 @@ impl ConversationRuntime {
|
||||
Some(pb::agent_client_message::Message::RunRequest(
|
||||
request,
|
||||
)) => {
|
||||
waiting_for_action = false;
|
||||
if draining {
|
||||
handle.reopen();
|
||||
draining = false;
|
||||
pending_finish = None;
|
||||
}
|
||||
if let Some(conversation_id) =
|
||||
request.conversation_id.as_deref()
|
||||
{
|
||||
@@ -117,60 +217,21 @@ impl ConversationRuntime {
|
||||
"invalid Cursor conversation id"
|
||||
);
|
||||
let _ = super::finish_failed(&handle, &error);
|
||||
let _ =
|
||||
handle.command(TransportCommand::Close).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
let previous_finished =
|
||||
if let Some(previous) = current.take() {
|
||||
previous.superseded.cancel();
|
||||
if let Some(run) = previous.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
for id in previous
|
||||
.tool_runtime
|
||||
.interrupt_for_run_replacement()
|
||||
.await
|
||||
{
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
Some(previous.finished.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (results, result_receiver) = tool_result_channel();
|
||||
let (runtime_actions, runtime_action_receiver) =
|
||||
mpsc::unbounded_channel::<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(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
dependencies.clone(),
|
||||
blob_sync.clone(),
|
||||
context_sync.clone(),
|
||||
generation,
|
||||
previous_finished,
|
||||
result_receiver,
|
||||
runtime_action_receiver,
|
||||
);
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
message,
|
||||
@@ -331,27 +392,48 @@ impl ConversationRuntime {
|
||||
// return an explicit Protocol Error rather than falling through silently.
|
||||
Some(
|
||||
pb::agent_client_message::Message::ConversationAction(
|
||||
action,
|
||||
conversation_action,
|
||||
),
|
||||
) => match action.action {
|
||||
) => match conversation_action.action.clone() {
|
||||
Some(
|
||||
pb::conversation_action::Action::UserMessageAction(
|
||||
action,
|
||||
),
|
||||
) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
let delivered_to_active_run =
|
||||
current.as_ref().is_some_and(|generation| {
|
||||
generation.run.lock().as_ref().is_some_and(
|
||||
|run| run.phase() == RunPhase::Running,
|
||||
) && generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(
|
||||
action.clone(),
|
||||
))
|
||||
.is_ok()
|
||||
});
|
||||
if delivered_to_active_run {
|
||||
continue;
|
||||
}
|
||||
let Some(previous) = current.as_ref() else {
|
||||
continue;
|
||||
};
|
||||
if generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(action))
|
||||
.is_err()
|
||||
{
|
||||
generation.results.send_error(crate::Error::Protocol(
|
||||
"UserMessageAction arrived without an active Run"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
let mut request = previous.request.clone();
|
||||
request.action = Some(conversation_action);
|
||||
request.conversation_state = None;
|
||||
request.pre_fetched_blobs.clear();
|
||||
waiting_for_action = false;
|
||||
start_generation(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||
if let Some(generation) = current.as_ref() {
|
||||
@@ -428,6 +510,95 @@ impl ConversationRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn start_generation(
|
||||
registry: &ConversationRegistry,
|
||||
handle: &TransportHandle,
|
||||
dependencies: &ConversationDependencies,
|
||||
blob_sync: &BlobSynchronizer,
|
||||
context_sync: &RequestContextSynchronizer,
|
||||
tool_runtime_factory: &CursorToolRuntime,
|
||||
current: &mut Option<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;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
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 {
|
||||
|
||||
@@ -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 inbox;
|
||||
mod lifecycle;
|
||||
mod output;
|
||||
mod registry;
|
||||
|
||||
pub use handle::*;
|
||||
pub use inbox::*;
|
||||
pub(crate) use lifecycle::*;
|
||||
pub use output::*;
|
||||
pub use registry::*;
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -422,6 +422,85 @@ async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() {
|
||||
assert_eq!(output.recv().await, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn queued_user_message_after_turn_ended_starts_the_next_turn() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(text_response("first turn"));
|
||||
provider.push(text_response("queued turn"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("queued-after-turn").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run_for(
|
||||
"queued-after-turn",
|
||||
"queued-after-turn-conversation",
|
||||
)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
wait_for_turn_ended(&handle, &mut output, &mut append_seqno).await;
|
||||
assert_transport_remains_open(&handle, &mut output, &mut append_seqno).await;
|
||||
|
||||
cursor_server::api::cursor::bidi::append(
|
||||
®istry,
|
||||
cursor_server::api::cursor::bidi::DecodedAppend {
|
||||
request_id: "queued-after-turn".into(),
|
||||
seqno: append_seqno,
|
||||
message: runtime_user_message(),
|
||||
},
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
|
||||
let mut text = String::new();
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("queued turn closed without EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
assert_eq!(payload.as_ref(), b"{}");
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message {
|
||||
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
|
||||
text.push_str(&delta.text);
|
||||
}
|
||||
}
|
||||
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
|
||||
}
|
||||
|
||||
assert!(text.contains("queued turn"));
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
assert_eq!(
|
||||
&requests[1].history[..requests[0].history.len()],
|
||||
requests[0].history.as_slice(),
|
||||
"queued continuation must preserve the first provider request as a prefix"
|
||||
);
|
||||
let history = serde_json::to_string(&requests[1].history).unwrap();
|
||||
assert!(history.contains("queued follow-up"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_user_message_action_interrupts_and_continues_with_new_message() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
@@ -1732,6 +1811,55 @@ async fn run_to_end(
|
||||
}
|
||||
}
|
||||
|
||||
async fn wait_for_turn_ended(
|
||||
handle: &cursor_server::cursor::TransportHandle,
|
||||
output: &mut tokio::sync::mpsc::UnboundedReceiver<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,
|
||||
|
||||
Reference in New Issue
Block a user