feat: initialize server structure and database schema

- Added initial server setup with Cargo.toml defining dependencies and project structure.
- Created build.rs for generating protobuf bindings and validating wire contracts.
- Established database schema with initial migration files for conversations, messages, and runs.
- Introduced tools and prompts for Cursor functionality, enhancing user interaction capabilities.
This commit is contained in:
leookun
2026-08-30 01:33:39 +08:00
parent d200b3791d
commit 44e2d8057a
206 changed files with 37923 additions and 21 deletions
+207
View File
@@ -0,0 +1,207 @@
//! Accepts ordered Cursor Bidi append requests and routes them by request_id.
use prost::Message;
use crate::{
cursor::{
conversation::TransportCommand,
protocol::{
events,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
transport::{TransportParent, TransportRegistry},
},
Error, Result,
};
pub struct DecodedAppend {
pub request_id: String,
pub seqno: i64,
pub message: agent::AgentClientMessage,
}
impl DecodedAppend {
pub fn model_id(&self) -> Option<&str> {
let agent::agent_client_message::Message::RunRequest(request) =
self.message.message.as_ref()?
else {
return None;
};
request
.requested_model
.as_ref()
.map(|model| model.model_id.as_str())
.filter(|model| !model.is_empty())
.or_else(|| {
request
.model_details
.as_ref()
.map(|model| model.model_id.as_str())
.filter(|model| !model.is_empty())
})
}
pub fn conversation_id(&self) -> Option<&str> {
let agent::agent_client_message::Message::RunRequest(request) =
self.message.message.as_ref()?
else {
return None;
};
request.conversation_id.as_deref()
}
pub fn is_background_task_completion(&self) -> bool {
let Some(agent::agent_client_message::Message::RunRequest(request)) =
self.message.message.as_ref()
else {
return false;
};
matches!(
request
.action
.as_ref()
.and_then(|action| action.action.as_ref()),
Some(agent::conversation_action::Action::BackgroundTaskCompletionAction(_))
)
}
pub fn trace_metadata(&self) -> serde_json::Value {
let Some(message) = self.message.message.as_ref() else {
return serde_json::json!({
"append_seqno": self.seqno,
"message_type": "empty",
});
};
let agent::agent_client_message::Message::RunRequest(request) = message else {
return serde_json::json!({
"append_seqno": self.seqno,
"message_type": client_message_type(message),
});
};
let (action_type, history_messages, history_images) = request
.action
.as_ref()
.and_then(|action| action.action.as_ref())
.map(|action| match action {
agent::conversation_action::Action::UserMessageAction(action) => {
let history = action.conversation_history.as_ref();
(
"user_message",
history.map_or(0, |history| history.messages.len()),
history.map_or(0, history_image_count),
)
}
agent::conversation_action::Action::BackgroundTaskCompletionAction(_) => {
("background_task_completion", 0, 0)
}
agent::conversation_action::Action::ExecutePlanAction(_) => ("execute_plan", 0, 0),
agent::conversation_action::Action::SummarizeAction(_) => ("summarize", 0, 0),
_ => ("other", 0, 0),
})
.unwrap_or(("none", 0, 0));
let state = request.conversation_state.as_ref();
serde_json::json!({
"append_seqno": self.seqno,
"message_type": "run_request",
"conversation_id": request.conversation_id,
"model_id": self.model_id(),
"action_type": action_type,
"conversation_history_messages": history_messages,
"conversation_history_images": history_images,
"root_message_count": state.map_or(0, |state| state.root_prompt_messages_json.len()),
"turn_count": state.map_or(0, |state| state.turns.len()),
"prefetched_blob_count": request.pre_fetched_blobs.len(),
})
}
}
fn client_message_type(message: &agent::agent_client_message::Message) -> &'static str {
use agent::agent_client_message::Message;
match message {
Message::RunRequest(_) => "run_request",
Message::ExecClientMessage(_) => "exec_client_message",
Message::ExecClientControlMessage(_) => "exec_client_control_message",
Message::KvClientMessage(_) => "kv_client_message",
Message::ConversationAction(_) => "conversation_action",
Message::InteractionResponse(_) => "interaction_response",
Message::ClientHeartbeat(_) => "client_heartbeat",
Message::PrewarmRequest(_) => "prewarm_request",
}
}
fn history_image_count(history: &agent::ConversationHistory) -> usize {
use agent::{
conversation_history_message::Message,
conversation_history_tool_result_content::Content as ToolContent,
conversation_history_user_content::Content as UserContent,
};
history
.messages
.iter()
.map(|message| match message.message.as_ref() {
Some(Message::User(user)) => user
.content
.iter()
.filter(|content| matches!(content.content, Some(UserContent::Image(_))))
.count(),
Some(Message::Tool(tool)) => tool
.content
.iter()
.filter(|content| matches!(content.content, Some(ToolContent::Image(_))))
.count(),
_ => 0,
})
.sum()
}
pub fn decode(request: &ai::BidiAppendRequest) -> Result<DecodedAppend> {
let request_id = request
.request_id
.as_ref()
.map(|id| id.request_id.as_str())
.filter(|id| !id.is_empty())
.ok_or_else(|| Error::Protocol("BidiAppend request_id is required".into()))?;
if !request.data_binary.is_empty() {
return Err(Error::Protocol(
"BidiAppend data_binary is not part of the captured protocol".into(),
));
}
if request.data.is_empty() {
return Err(Error::Protocol(
"BidiAppend contains no AgentClientMessage".into(),
));
}
let payload = hex::decode(&request.data)
.map_err(|error| Error::Protocol(format!("invalid BidiAppend hex: {error}")))?;
Ok(DecodedAppend {
request_id: request_id.into(),
seqno: request.append_seqno,
message: agent::AgentClientMessage::decode(payload.as_slice())?,
})
}
pub async fn append(
registry: &TransportRegistry,
request: DecodedAppend,
parent: Option<TransportParent>,
) -> Result<ai::BidiAppendResponse> {
let handle = registry.get_or_create(&request.request_id).await?;
if let Some(conversation_id) = request.conversation_id() {
handle.set_conversation_id(conversation_id)?;
}
if let Some(parent) = parent {
handle.set_parent(parent)?;
}
if matches!(
request.message.message.as_ref(),
Some(agent::agent_client_message::Message::ClientHeartbeat(_))
) {
handle.emit(&events::heartbeat())?;
}
handle
.command(TransportCommand::Append {
seqno: request.seqno,
message: Box::new(request.message),
})
.await?;
Ok(ai::BidiAppendResponse {})
}
+227
View File
@@ -0,0 +1,227 @@
//! Implements Cursor HTTP endpoints outside the Agent Run stream.
use axum::{
body::{to_bytes, Body, Bytes},
extract::{DefaultBodyLimit, Extension, State},
http::{header, HeaderMap, HeaderValue, Request, Response, StatusCode},
routing::{get, post},
Router,
};
use tower_http::decompression::RequestDecompressionLayer;
use crate::{
api::cursor::{
bidi,
proxy::{self, CursorProxy},
run_sse,
},
cursor::{
protocol::{
connect,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
transport::{TransportParent, TransportRegistry},
},
Result,
};
pub fn router(registry: TransportRegistry) -> Result<Router> {
let proxy = CursorProxy::cursor(registry.store().clone())?;
Ok(router_with_proxy(registry, proxy))
}
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
Router::new()
.route("/__byok-api__/healthz", get(health))
.route("/agent.v1.AgentService/RunSSE", post(run_sse_handler))
.route("/aiserver.v1.BidiService/BidiAppend", post(bidi_handler))
.route(
"/aiserver.v1.AiService/AvailableModels",
post(model_catalog::available_models),
)
.route(
"/agent.v1.AgentService/GetUsableModels",
post(model_catalog::usable_models),
)
.route(
"/aiserver.v1.AiService/GetUsableModels",
post(model_catalog::usable_models),
)
.route(
"/aiserver.v1.AuthService/GetEmail",
post(account::get_email),
)
.route("/aiserver.v1.DashboardService/GetMe", post(account::get_me))
.route(
"/aiserver.v1.DashboardService/GetTeams",
post(account::get_teams),
)
.route(
"/aiserver.v1.DashboardService/GetUserProfile",
post(account::get_user_profile),
)
.route(
"/aiserver.v1.DashboardService/GetCurrentPeriodUsage",
post(account::current_period_usage),
)
.route(
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
post(account::usage_limit_status),
)
.route(
analytics::BOOTSTRAP_STATSIG_PATH,
post(analytics::bootstrap_statsig),
)
.route("/auth/full_stripe_profile", get(account::stripe_profile))
.merge(tab::router())
.route_layer(DefaultBodyLimit::disable())
.route_layer(RequestDecompressionLayer::new())
.fallback(proxy::forward)
.method_not_allowed_fallback(proxy::forward)
.layer(Extension(proxy))
.with_state(registry)
}
async fn health() -> StatusCode {
StatusCode::NO_CONTENT
}
async fn run_sse_handler(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
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;
}
match route {
crate::cursor::transport::TransportRoute::Local => {
run_sse::stream(&registry, &request.request_id).await
}
crate::cursor::transport::TransportRoute::Upstream(generation) => {
let response = proxy::forward(
Extension(proxy),
Request::from_parts(parts, Body::from(body)),
)
.await?;
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
}
}
}
async fn bidi_handler(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let request: ai::BidiAppendRequest = connect::decode_unary(&body)?;
let decoded = bidi::decode(&request)?;
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 local = if let Some(model_id) = decoded.model_id() {
if registry.store().model(model_id).await?.is_some() {
tracing::info!(
request_id = decoded.request_id,
model_id,
"routing Cursor Run to BYOK provider"
);
true
} else {
tracing::info!(
request_id = decoded.request_id,
model_id,
"routing Cursor Run to Cursor upstream"
);
false
}
} else if registry.local(&decoded.request_id).await.is_some() {
true
} else if registry.upstream(&decoded.request_id).await {
false
} else {
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,
conversation_id.as_deref(),
if local {
"local_byok"
} else {
"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;
}
if !local {
if first_model.is_some() {
registry.mark_upstream(&decoded.request_id).await;
}
return proxy::forward(
Extension(proxy),
Request::from_parts(parts, Body::from(body)),
)
.await;
}
let parent = parent_headers(&parts.headers)?;
bidi::append(&registry, decoded, parent).await?;
let mut response = Response::new(axum::body::Body::empty());
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
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)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok((parts, body))
}
fn parent_headers(headers: &HeaderMap) -> Result<Option<TransportParent>> {
let request_id = header_text(headers, "x-parent-request-id")?;
let tool_call_id = header_text(headers, "x-parent-agent-tool-call-id")?;
match (request_id, tool_call_id) {
(None, None) => Ok(None),
(Some(request_id), Some(tool_call_id)) => Ok(Some(TransportParent {
request_id: request_id.into(),
tool_call_id: tool_call_id.into(),
})),
_ => Err(crate::Error::Protocol(
"Cursor subagent request must include both parent headers".into(),
)),
}
}
fn header_text<'a>(headers: &'a HeaderMap, name: &str) -> Result<Option<&'a str>> {
headers
.get(name)
.map(|value| value.to_str())
.transpose()
.map_err(|error| crate::Error::Protocol(format!("invalid {name} header: {error}")))
}
+8
View File
@@ -0,0 +1,8 @@
//! Wires the Cursor-facing API routes.
pub mod bidi;
mod handlers;
pub mod proxy;
mod run_sse;
pub use handlers::router;
+236
View File
@@ -0,0 +1,236 @@
//! Selects local handling or the configured official Cursor upstream.
use std::time::Instant;
use axum::{
body::{to_bytes, Body, Bytes},
extract::Extension,
http::{header, Request, Response},
};
use crate::Result;
const CURSOR_UPSTREAM: &str = "https://api2.cursor.sh";
pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url";
#[derive(Clone)]
pub struct CursorProxy {
client: Option<reqwest::Client>,
store: Option<crate::store::Store>,
upstream: String,
}
pub struct BufferedResponse {
pub status: axum::http::StatusCode,
pub headers: axum::http::HeaderMap,
pub body: Bytes,
}
impl BufferedResponse {
pub fn into_response(self) -> Response<Body> {
let body = self.body.clone();
self.with_body(body)
}
pub fn with_body(mut self, body: Bytes) -> Response<Body> {
self.headers.insert(
header::CONTENT_LENGTH,
body.len()
.to_string()
.parse()
.expect("body length is always a valid header value"),
);
let mut response = Response::new(Body::from(body));
*response.status_mut() = self.status;
*response.headers_mut() = self.headers;
response
}
}
impl CursorProxy {
pub fn cursor(store: crate::store::Store) -> Result<Self> {
Ok(Self {
client: None,
store: Some(store),
upstream: CURSOR_UPSTREAM.into(),
})
}
async fn client(&self) -> Result<reqwest::Client> {
match (&self.client, &self.store) {
(Some(client), _) => Ok(client.clone()),
(_, Some(store)) => Ok(crate::network::client_builder(store)
.await?
.redirect(reqwest::redirect::Policy::none())
.build()?),
_ => unreachable!("Cursor proxy always has a client or store"),
}
}
}
pub async fn forward(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_request(&proxy, request, None).await
}
pub(crate) async fn forward_to_service(
proxy: &CursorProxy,
request: Request<Body>,
service_url: &str,
) -> Result<Response<Body>> {
forward_request(proxy, request, Some(service_url)).await
}
async fn forward_request(
proxy: &CursorProxy,
request: Request<Body>,
service_url: Option<&str>,
) -> Result<Response<Body>> {
let started = Instant::now();
let (parts, body) = request.into_parts();
let path = parts
.uri
.path_and_query()
.map_or("/", |value| value.as_str())
.to_owned();
let url = match service_url {
Some(service_url) => format!("{}{}", service_url.trim_end_matches('/'), path),
None => upstream_url(&parts.headers, &proxy.upstream, &path)?,
};
let mut headers = parts.headers;
headers.remove(UPSTREAM_URL_HEADER);
headers.remove(header::HOST);
remove_hop_by_hop_headers(&mut headers);
let client = proxy.client().await?;
let upstream = client
.request(parts.method.clone(), url)
.headers(headers)
.body(reqwest::Body::wrap_stream(body.into_data_stream()))
.send()
.await;
let upstream = match upstream {
Ok(response) => response,
Err(error) => {
tracing::error!(
method = %parts.method,
path,
elapsed_ms = started.elapsed().as_millis(),
%error,
"Cursor upstream request failed"
);
return Err(error.into());
}
};
let status = upstream.status();
let mut response_headers = upstream.headers().clone();
remove_hop_by_hop_headers(&mut response_headers);
let mut response = Response::new(Body::from_stream(upstream.bytes_stream()));
*response.status_mut() = status;
*response.headers_mut() = response_headers;
tracing::info!(
method = %parts.method,
path,
%status,
elapsed_ms = started.elapsed().as_millis(),
"forwarded Cursor backend request"
);
Ok(response)
}
pub async fn forward_buffered(
proxy: &CursorProxy,
request: Request<Body>,
) -> Result<BufferedResponse> {
let (parts, body) = request.into_parts();
let path = parts
.uri
.path_and_query()
.map_or("/", |value| value.as_str());
let url = upstream_url(&parts.headers, &proxy.upstream, path)?;
let mut headers = parts.headers;
headers.remove(UPSTREAM_URL_HEADER);
headers.remove(header::HOST);
remove_hop_by_hop_headers(&mut headers);
headers.insert(
"connect-accept-encoding",
axum::http::HeaderValue::from_static("identity"),
);
headers.insert(
header::ACCEPT_ENCODING,
axum::http::HeaderValue::from_static("identity"),
);
let body = to_bytes(body, usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
let upstream = proxy
.client()
.await?
.request(parts.method, url)
.headers(headers)
.body(body)
.send()
.await?;
let status = upstream.status();
let mut headers = upstream.headers().clone();
remove_hop_by_hop_headers(&mut headers);
let body = upstream.bytes().await?;
Ok(BufferedResponse {
status,
headers,
body,
})
}
fn upstream_url(headers: &axum::http::HeaderMap, fallback: &str, path: &str) -> Result<String> {
let Some(value) = headers.get(UPSTREAM_URL_HEADER) else {
return Ok(format!("{fallback}{path}"));
};
let value = value
.to_str()
.map_err(|error| crate::Error::Protocol(format!("invalid upstream URL header: {error}")))?;
let url = reqwest::Url::parse(value)
.map_err(|error| crate::Error::Protocol(format!("invalid upstream URL: {error}")))?;
let host = url.host_str().unwrap_or_default();
if url.scheme() != "https" || !crate::local_app::proxy_host_allowed(host) {
return Err(crate::Error::Protocol(
"upstream URL must target a Cursor HTTPS host".into(),
));
}
Ok(url.into())
}
fn remove_hop_by_hop_headers(headers: &mut axum::http::HeaderMap) {
let connection_headers = headers
.get(header::CONNECTION)
.and_then(|value| value.to_str().ok())
.map(|value| {
value
.split(',')
.map(str::trim)
.filter(|name| !name.is_empty())
.map(str::to_owned)
.collect::<Vec<_>>()
})
.unwrap_or_default();
for name in connection_headers {
headers.remove(name);
}
for name in [
header::CONNECTION,
header::PROXY_AUTHENTICATE,
header::PROXY_AUTHORIZATION,
header::TE,
header::TRAILER,
header::TRANSFER_ENCODING,
header::UPGRADE,
] {
headers.remove(name);
}
headers.remove("keep-alive");
}
+232
View File
@@ -0,0 +1,232 @@
//! Subscribes Cursor RunSSE clients to replayable Transport output.
use axum::{
body::Body,
http::{header, HeaderValue, Response, StatusCode},
};
use bytes::Bytes;
use std::convert::Infallible;
use tokio::sync::mpsc;
use tokio_stream::StreamExt;
use crate::{
cursor::{
protocol::connect::{self, END_STREAM_FLAG},
services::observability::CursorTraceRecorder,
transport::{TransportHandle, TransportRegistry},
},
Result,
};
pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Response<Body>> {
let handle = registry.get_or_create(request_id).await?;
let receiver = handle.subscribe();
let trace = handle.trace().cloned();
if let Some(trace) = &trace {
trace.response_started(StatusCode::OK.as_u16()).await;
}
let body_stream = local_body_stream(receiver, handle, trace);
let mut response = Response::new(Body::from_stream(body_stream));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("text/event-stream"),
);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-cache"));
response
.headers_mut()
.insert("connect-protocol-version", HeaderValue::from_static("1"));
Ok(response)
}
fn local_body_stream(
mut receiver: mpsc::UnboundedReceiver<Bytes>,
handle: TransportHandle,
trace: Option<CursorTraceRecorder>,
) -> impl tokio_stream::Stream<Item = std::result::Result<Bytes, Infallible>> {
async_stream::stream! {
let mut guard = LocalRunGuard::new(handle);
let mut trace = TraceStreamSink::new(trace, "byok_server");
while let Some(chunk) = receiver.recv().await {
let terminal = is_end_stream_frame(&chunk);
trace.chunk(&chunk);
if terminal {
guard.complete();
trace.finish(end_stream_error(&chunk));
}
yield Ok::<Bytes, Infallible>(chunk);
if terminal {
return;
}
}
guard.complete();
trace.finish(None);
}
}
fn is_end_stream_frame(frame: &Bytes) -> bool {
frame
.first()
.is_some_and(|flags| flags & END_STREAM_FLAG != 0)
}
fn end_stream_error(frame: &Bytes) -> Option<String> {
connect::decode_frames(frame)
.ok()?
.into_iter()
.find_map(|(flags, payload)| {
if flags & END_STREAM_FLAG == 0 {
return None;
}
let value = serde_json::from_slice::<serde_json::Value>(&payload).ok()?;
let error = value.get("error")?;
let code = error.get("code").and_then(serde_json::Value::as_str);
let message = error
.get("message")
.and_then(serde_json::Value::as_str)
.filter(|message| !message.is_empty());
Some(match (code, message) {
(Some(code), Some(message)) => format!("{code}: {message}"),
(Some(code), None) => code.to_string(),
(None, Some(message)) => message.to_string(),
(None, None) => error.to_string(),
})
})
}
struct LocalRunGuard {
handle: TransportHandle,
completed: bool,
}
impl LocalRunGuard {
fn new(handle: TransportHandle) -> Self {
Self {
handle,
completed: false,
}
}
fn complete(&mut self) {
self.completed = true;
}
}
impl Drop for LocalRunGuard {
fn drop(&mut self) {
if !self.completed {
let handle = self.handle.clone();
tokio::spawn(async move {
handle.disconnect().await;
});
}
}
}
pub async fn upstream(
registry: TransportRegistry,
request_id: String,
generation: u64,
response: Response<Body>,
trace: Option<CursorTraceRecorder>,
) -> Response<Body> {
let (parts, body) = response.into_parts();
if let Some(trace) = &trace {
trace.response_started(parts.status.as_u16()).await;
}
let stream = async_stream::stream! {
let _guard = UpstreamRunGuard {
registry,
request_id,
generation,
};
let mut trace = TraceStreamSink::new(trace, "cursor_official");
let mut body = body.into_data_stream();
while let Some(chunk) = body.next().await {
match chunk {
Ok(chunk) => {
trace.chunk(&chunk);
yield Ok::<Bytes, axum::Error>(chunk);
}
Err(error) => {
trace.finish(Some(error.to_string()));
yield Err(error);
return;
}
}
}
trace.finish(None);
};
Response::from_parts(parts, Body::from_stream(stream))
}
enum TraceStreamEvent {
Chunk(Bytes),
Finish(Option<String>),
}
struct TraceStreamSink {
sender: Option<mpsc::UnboundedSender<TraceStreamEvent>>,
}
impl TraceStreamSink {
fn new(trace: Option<CursorTraceRecorder>, source: &'static str) -> Self {
let Some(trace) = trace else {
return Self { sender: None };
};
let (sender, mut receiver) = mpsc::unbounded_channel();
tokio::spawn(async move {
while let Some(event) = receiver.recv().await {
match event {
TraceStreamEvent::Chunk(chunk) => {
trace.response_chunk(source, &chunk).await;
}
TraceStreamEvent::Finish(error) => {
trace.finish(error.as_deref()).await;
return;
}
}
}
trace.finish(None).await;
});
Self {
sender: Some(sender),
}
}
fn chunk(&self, chunk: &Bytes) {
if let Some(sender) = &self.sender {
let _ = sender.send(TraceStreamEvent::Chunk(chunk.clone()));
}
}
fn finish(&mut self, error: Option<String>) {
if let Some(sender) = self.sender.take() {
let _ = sender.send(TraceStreamEvent::Finish(error));
}
}
}
impl Drop for TraceStreamSink {
fn drop(&mut self) {
if self.sender.is_some() {
self.finish(Some(
"response stream dropped before completion".to_string(),
));
}
}
}
struct UpstreamRunGuard {
registry: TransportRegistry,
request_id: String,
generation: u64,
}
impl Drop for UpstreamRunGuard {
fn drop(&mut self) {
self.registry
.finish_upstream(self.request_id.clone(), self.generation);
}
}
+6
View File
@@ -0,0 +1,6 @@
//! Exposes the HTTP and Connect API layer.
pub mod cursor;
mod router;
pub use router::router;
+7
View File
@@ -0,0 +1,7 @@
//! Builds the top-level server router.
use crate::{cursor::transport::TransportRegistry, Result};
pub fn router(registry: TransportRegistry) -> Result<axum::Router> {
super::cursor::router(registry)
}
+170
View File
@@ -0,0 +1,170 @@
//! Assembles server dependencies and starts the application services.
use std::{future::IntoFuture, net::SocketAddr, time::Duration};
use tokio::net::TcpListener;
use tokio_util::sync::CancellationToken;
use crate::{
api,
config::{Config, ConsoleSource},
control,
cursor::{
prompting::{PromptAssets, PromptCompiler},
transport::TransportRegistry,
},
local_app::CursorHarness,
provider::ProviderRouter,
store::Store,
Result,
};
pub struct App {
config: Config,
router: axum::Router,
registry: TransportRegistry,
harness: CursorHarness,
store: Store,
}
impl App {
pub async fn new(mut config: Config) -> Result<Self> {
let store = Store::connect(&config.database_url).await?;
if config.use_persisted_ports {
config
.listen_addr
.set_port(store.port_settings().await?.service_port);
}
let assets = PromptAssets::embedded()?;
let compiler = PromptCompiler::new(assets);
let provider = std::sync::Arc::new(ProviderRouter::new(
store.clone(),
config.provider_request_timeout,
));
let registry = TransportRegistry::new(store.clone(), provider.clone(), compiler);
let control = control::ControlService::new(store.clone(), provider)?;
let harness = control.cursor_harness().clone();
let mut router = api::router(registry.clone())?;
router = match &config.console {
Some(ConsoleSource::Directory(directory)) => {
router.merge(control::web_router(control.clone(), directory))
}
Some(ConsoleSource::Proxy(target)) => {
router.merge(control::proxy_web_router(control.clone(), target.clone()))
}
None => router.merge(control::api_router(control.clone())),
};
Ok(Self {
router,
registry,
harness,
store,
config,
})
}
pub fn merge_router(mut self, router: axum::Router) -> Self {
self.router = self.router.merge(router);
self
}
pub async fn bind(&self) -> Result<TcpListener> {
let requested = self.config.listen_addr;
let listener = bind_service_listener(requested, self.config.use_persisted_ports).await?;
if self.config.use_persisted_ports {
self.store
.set_service_port(listener.local_addr()?.port())
.await?;
}
Ok(listener)
}
pub fn harness(&self) -> CursorHarness {
self.harness.clone()
}
pub fn store(&self) -> Store {
self.store.clone()
}
pub async fn serve(self) -> Result<()> {
let listener = self.bind().await?;
let shutdown = CancellationToken::new();
let signal_shutdown = shutdown.clone();
let running = self.serve_on(listener, shutdown);
tokio::pin!(running);
tokio::select! {
result = &mut running => result,
() = shutdown_signal() => {
tracing::info!("shutdown signal received; cancelling active runs");
signal_shutdown.cancel();
running.await
}
}
}
pub async fn serve_on(self, listener: TcpListener, shutdown: CancellationToken) -> Result<()> {
let address = listener.local_addr()?;
self.harness.set_backend_addr(address);
tracing::info!(%address, "cursor server listening");
let registry = self.registry;
let harness = self.harness;
let graceful = shutdown.clone();
let server = axum::serve(listener, self.router)
.with_graceful_shutdown(async move {
graceful.cancelled().await;
})
.into_future();
tokio::pin!(server);
tokio::select! {
result = &mut server => {
if let Err(error) = harness.disable().await {
tracing::warn!(%error, "failed to disable Cursor harness after server stop");
}
result?
},
() = shutdown.cancelled() => {
if let Err(error) = harness.disable().await {
tracing::warn!(%error, "failed to disable Cursor harness during shutdown");
}
registry.shutdown().await;
match tokio::time::timeout(Duration::from_secs(10), &mut server).await {
Ok(result) => result?,
Err(_) => tracing::warn!("graceful shutdown timed out; forcing server close"),
}
}
}
Ok(())
}
}
async fn bind_service_listener(
requested: SocketAddr,
allow_random_fallback: bool,
) -> Result<TcpListener> {
match TcpListener::bind(requested).await {
Ok(listener) => Ok(listener),
Err(error) if allow_random_fallback && requested.port() != 0 => {
tracing::warn!(%requested, %error, "configured service port unavailable; selecting a random port");
Ok(TcpListener::bind(SocketAddr::new(requested.ip(), 0)).await?)
}
Err(error) => Err(error.into()),
}
}
async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
let terminate = async {
if let Ok(mut signal) =
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
{
signal.recv().await;
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! { _ = ctrl_c => {}, _ = terminate => {} }
}
+16
View File
@@ -0,0 +1,16 @@
//! Starts the Cursor BYOK server executable.
use cursor_server::{App, Config, Result};
use tracing_subscriber::prelude::*;
#[tokio::main]
async fn main() -> Result<()> {
tracing_subscriber::registry()
.with(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "cursor_server=info".into()),
)
.with(tracing_subscriber::fmt::layer())
.init();
App::new(Config::from_env()?).await?.serve().await
}
+144
View File
@@ -0,0 +1,144 @@
//! Loads and validates process-level server configuration.
use std::{env, fs, net::SocketAddr, path::PathBuf, time::Duration};
#[cfg(unix)]
use std::os::unix::fs::PermissionsExt;
use crate::{Error, Result};
const DATA_DIR_NAME: &str = ".cursor-byok-v3";
const DATABASE_FILE_NAME: &str = "cursor-byok.db";
const V0049_DATA_DIR_NAME: &str = ".cursor-local-assistant-v2";
const V0049_CONFIG_FILE_NAME: &str = "config.yaml";
const DEFAULT_PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(3000);
pub fn managed_data_dir() -> Result<PathBuf> {
let home_dir = dirs::home_dir()
.ok_or_else(|| Error::Config("cannot resolve user home directory".into()))?;
let data_dir = home_dir.join(DATA_DIR_NAME);
fs::create_dir_all(&data_dir)?;
#[cfg(unix)]
fs::set_permissions(&data_dir, fs::Permissions::from_mode(0o700))?;
Ok(data_dir)
}
pub fn v0049_config_path() -> Result<PathBuf> {
let home_dir = dirs::home_dir()
.ok_or_else(|| Error::Config("cannot resolve user home directory".into()))?;
Ok(home_dir
.join(V0049_DATA_DIR_NAME)
.join(V0049_CONFIG_FILE_NAME))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ProviderKind {
OpenAiChat,
OpenAiResponses,
Anthropic,
}
#[derive(Clone)]
pub struct ProviderConfig {
pub kind: ProviderKind,
pub request_url: String,
pub api_key: String,
pub custom_headers: reqwest::header::HeaderMap,
pub max_output_tokens: Option<u64>,
pub request_timeout: Duration,
}
#[derive(Clone)]
pub struct Config {
pub listen_addr: SocketAddr,
pub database_url: String,
pub provider_request_timeout: Duration,
pub console: Option<ConsoleSource>,
pub use_persisted_ports: bool,
}
#[derive(Clone)]
pub enum ConsoleSource {
Directory(PathBuf),
Proxy(url::Url),
}
impl Config {
pub fn from_env() -> Result<Self> {
let listen_addr = env::var("CURSOR_LISTEN_ADDR")
.unwrap_or_else(|_| "127.0.0.1:3000".into())
.parse()
.map_err(|error| Error::Config(format!("invalid CURSOR_LISTEN_ADDR: {error}")))?;
let request_timeout = match env::var("CURSOR_PROVIDER_TIMEOUT_SECONDS") {
Ok(value) => Duration::from_secs(value.parse().map_err(|error| {
Error::Config(format!("invalid CURSOR_PROVIDER_TIMEOUT_SECONDS: {error}"))
})?),
Err(env::VarError::NotPresent) => DEFAULT_PROVIDER_REQUEST_TIMEOUT,
Err(error) => {
return Err(Error::Config(format!(
"invalid CURSOR_PROVIDER_TIMEOUT_SECONDS: {error}"
)))
}
};
let console_dir = env::var_os("CURSOR_CONSOLE_DIR").map(PathBuf::from);
let console_proxy = env::var("CURSOR_CONSOLE_PROXY")
.ok()
.map(|value| {
value.parse().map_err(|error| {
Error::Config(format!("invalid CURSOR_CONSOLE_PROXY: {error}"))
})
})
.transpose()?;
let console = match (console_dir, console_proxy) {
(Some(_), Some(_)) => {
return Err(Error::Config(
"CURSOR_CONSOLE_DIR and CURSOR_CONSOLE_PROXY cannot both be set".into(),
))
}
(Some(directory), None) => Some(ConsoleSource::Directory(directory)),
(None, Some(proxy)) => Some(ConsoleSource::Proxy(proxy)),
(None, None) => None,
};
Ok(Self {
listen_addr,
database_url: database_url_from_env()?,
provider_request_timeout: request_timeout,
console,
use_persisted_ports: false,
})
}
pub fn desktop() -> Result<Self> {
Ok(Self {
listen_addr: "127.0.0.1:0"
.parse()
.expect("desktop listen address is static"),
database_url: default_database_url()?,
provider_request_timeout: DEFAULT_PROVIDER_REQUEST_TIMEOUT,
console: None,
use_persisted_ports: true,
})
}
}
fn database_url_from_env() -> Result<String> {
match env::var("CURSOR_DATABASE_URL") {
Ok(database_url) => Ok(database_url),
Err(env::VarError::NotPresent) => default_database_url(),
Err(error) => Err(Error::Config(format!(
"invalid CURSOR_DATABASE_URL: {error}"
))),
}
}
fn default_database_url() -> Result<String> {
let data_dir = managed_data_dir()?;
database_url_for_dir(&data_dir)
}
fn database_url_for_dir(data_dir: &std::path::Path) -> Result<String> {
let database_path = data_dir.join(DATABASE_FILE_NAME);
let database_path = database_path
.to_str()
.ok_or_else(|| Error::Config("database path is not valid UTF-8".into()))?;
Ok(format!("sqlite://{database_path}"))
}
+148
View File
@@ -0,0 +1,148 @@
//! Implements advertisement configuration endpoints.
//! Advertisement service contract and desktop HTTP handler.
use axum::{
extract::{Path, State},
http::{HeaderMap, StatusCode},
Json,
};
use serde::{Deserialize, Serialize};
use url::Url;
use crate::{Error, Result};
use super::ControlService;
// 此广告拉取不涉及用户隐私,用户id随机产生
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/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";
pub(super) const DISABLED_AD_IDS_HEADER: &str = "disable-ad-ids";
pub(super) const LANGUAGE_HEADER: &str = "accept-language";
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdRuntime {
pub slots: Vec<AdSlot>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct AdSlot {
pub id: String,
pub enabled: bool,
pub placement: AdPlacement,
pub target: AdTarget,
pub content: AdContent,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum AdPlacement {
Menu,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct AdTarget {
pub title: String,
pub description: String,
pub image_url: String,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct AdContent {
pub title: String,
pub description: String,
pub image_url: String,
pub details: Vec<AdDetail>,
pub button: AdButton,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdDetail {
pub label: String,
pub value: String,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdButton {
pub label: String,
pub action: AdAction,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdAction {
#[serde(rename = "type")]
pub action_type: AdActionType,
pub url: String,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdDismissalInput {
pub reason: String,
}
#[derive(Clone, Copy, Debug, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum AdActionType {
OpenBrowser,
}
impl AdRuntime {
pub(super) fn into_menu_slots(mut self) -> Result<Self> {
self.slots
.retain(|slot| slot.enabled && slot.placement == AdPlacement::Menu);
for slot in &self.slots {
validate_http_url(&slot.target.image_url, "target.imageUrl")?;
validate_http_url(&slot.content.image_url, "content.imageUrl")?;
validate_http_url(&slot.content.button.action.url, "content.button.action.url")?;
}
Ok(self)
}
}
fn validate_http_url(value: &str, field: &str) -> Result<()> {
let url = Url::parse(value)
.map_err(|error| Error::Provider(format!("advertisement {field} is invalid: {error}")))?;
if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() {
return Err(Error::Provider(format!(
"advertisement {field} must be an absolute HTTP or HTTPS URL"
)));
}
Ok(())
}
pub async fn get(
State(service): State<ControlService>,
headers: HeaderMap,
) -> Result<Json<AdRuntime>> {
let disabled_ad_ids = headers
.get(DISABLED_AD_IDS_HEADER)
.and_then(|value| value.to_str().ok());
Ok(Json(
service.ads(disabled_ad_ids, ad_language(&headers)).await?,
))
}
fn ad_language(headers: &HeaderMap) -> &'static str {
match headers
.get(LANGUAGE_HEADER)
.and_then(|value| value.to_str().ok())
{
Some(value) if value.eq_ignore_ascii_case("zh-CN") => "zh-CN",
_ => "en-US",
}
}
pub async fn dismiss(
State(service): State<ControlService>,
Path(ad_id): Path<String>,
Json(input): Json<AdDismissalInput>,
) -> Result<StatusCode> {
service.dismiss_ad(&ad_id, &input).await?;
Ok(StatusCode::NO_CONTENT)
}
+34
View File
@@ -0,0 +1,34 @@
//! Implements provider call inspection endpoints.
use axum::{
extract::{Path, Query, State},
Json,
};
use serde::Deserialize;
use crate::Result;
use super::{CallDetail, CallSummary, ControlService};
#[derive(Deserialize)]
pub struct CallQuery {
#[serde(default = "default_limit")]
limit: i64,
}
pub async fn list(
State(service): State<ControlService>,
Query(query): Query<CallQuery>,
) -> Result<Json<Vec<CallSummary>>> {
Ok(Json(service.calls(query.limit).await?))
}
pub async fn detail(
State(service): State<ControlService>,
Path(call_id): Path<String>,
) -> Result<Json<CallDetail>> {
Ok(Json(service.call(&call_id).await?))
}
fn default_limit() -> i64 {
100
}
+28
View File
@@ -0,0 +1,28 @@
//! Implements local application control endpoints.
use axum::{extract::State, Json};
use crate::{
local_app::{CursorHarnessStatus, SetEnabled},
Result,
};
use super::ControlService;
pub async fn status(State(service): State<ControlService>) -> Result<Json<CursorHarnessStatus>> {
Ok(Json(service.cursor_harness().status().await?))
}
pub async fn initialize_ca(
State(service): State<ControlService>,
) -> Result<Json<CursorHarnessStatus>> {
Ok(Json(service.cursor_harness().initialize_ca().await?))
}
pub async fn set_enabled(
State(service): State<ControlService>,
Json(input): Json<SetEnabled>,
) -> Result<Json<CursorHarnessStatus>> {
Ok(Json(
service.cursor_harness().set_enabled(input.enabled).await?,
))
}
+221
View File
@@ -0,0 +1,221 @@
//! Exposes the local control API.
mod ads;
mod calls;
mod harness;
mod models;
mod overview;
mod service;
mod settings;
use axum::{
body::{to_bytes, Body},
extract::State,
http::{header, header::CONTENT_TYPE, HeaderValue, Method, Request, Response, StatusCode},
routing::{any, get, post, put},
Router,
};
use tower_http::{
cors::{AllowOrigin, CorsLayer},
services::ServeDir,
};
use url::{Host, Url};
pub use service::{
CallDetail, CallSummary, ControlService, DiscoveredModels, LegacyModelImportPreview,
LegacyModelImportResult, ModelConnectivityResult, ModelDiscoveryInput, ObservabilitySettings,
};
pub fn web_router(service: ControlService, assets: impl AsRef<std::path::Path>) -> Router {
Router::new()
.nest_service(
"/__byok-api__",
ServeDir::new(assets).append_index_html_on_directories(true),
)
.merge(api_router(service))
}
pub fn proxy_web_router(service: ControlService, target: Url) -> Router {
frontend_proxy_router(target).merge(api_router(service))
}
fn frontend_proxy_router(target: Url) -> Router {
let state = FrontendProxy {
client: reqwest::Client::new(),
target: target.as_str().trim_end_matches('/').to_string(),
};
Router::new()
.route("/__byok-api__/", any(proxy_frontend))
.route("/__byok-api__/{*path}", any(proxy_frontend))
.with_state(state)
}
#[derive(Clone)]
struct FrontendProxy {
client: reqwest::Client,
target: String,
}
async fn proxy_frontend(
State(proxy): State<FrontendProxy>,
request: Request<Body>,
) -> Response<Body> {
let (parts, body) = request.into_parts();
let path = parts
.uri
.path_and_query()
.map(|value| value.as_str())
.unwrap_or("/__byok-api__/");
let mut upstream = proxy
.client
.request(parts.method, format!("{}{path}", proxy.target));
for (name, value) in &parts.headers {
if name != header::HOST && name != header::CONNECTION {
upstream = upstream.header(name, value);
}
}
let body = match to_bytes(body, 64 * 1024 * 1024).await {
Ok(body) => body,
Err(error) => return proxy_error(error),
};
let upstream = match upstream.body(body).send().await {
Ok(response) => response,
Err(error) => return proxy_error(error),
};
let status = upstream.status();
let headers = upstream.headers().clone();
let body = match upstream.bytes().await {
Ok(body) => body,
Err(error) => return proxy_error(error),
};
let mut response = Response::new(Body::from(body));
*response.status_mut() = status;
for (name, value) in &headers {
if name != header::CONNECTION
&& name != header::TRANSFER_ENCODING
&& name != header::CONTENT_LENGTH
{
response.headers_mut().insert(name, value.clone());
}
}
response
}
fn proxy_error(error: impl std::fmt::Display) -> Response<Body> {
tracing::warn!(%error, "frontend development proxy failed");
Response::builder()
.status(StatusCode::BAD_GATEWAY)
.body(Body::from("frontend development server is unavailable"))
.expect("static proxy error response")
}
pub fn api_router(service: ControlService) -> Router {
Router::new()
.route("/__byok-api__/api/ads", get(ads::get))
.route(
"/__byok-api__/api/ads/{ad_id}/dismissals",
post(ads::dismiss),
)
.route(
"/__byok-api__/api/models",
get(models::list).post(models::create),
)
.route("/__byok-api__/api/models/discover", post(models::discover))
.route(
"/__byok-api__/api/models/import-v0049",
get(models::preview_v0049).post(models::import_v0049),
)
.route("/__byok-api__/api/models/order", put(models::reorder))
.route("/__byok-api__/api/overview", get(overview::get))
.route(
"/__byok-api__/api/models/{model_hash}",
put(models::update).delete(models::remove),
)
.route(
"/__byok-api__/api/models/{model_hash}/test/{test_id}",
post(models::test).delete(models::cancel),
)
.route("/__byok-api__/api/llm-calls", get(calls::list))
.route("/__byok-api__/api/llm-calls/{call_id}", get(calls::detail))
.route(
"/__byok-api__/api/settings/observability",
get(settings::get).put(settings::update),
)
.route(
"/__byok-api__/api/settings/ports",
get(settings::get_ports).put(settings::update_ports),
)
.route(
"/__byok-api__/api/settings/storage/statistics",
get(settings::get_storage).delete(settings::clear_storage),
)
.route(
"/__byok-api__/api/settings/proxy",
get(settings::get_proxy).put(settings::update_proxy),
)
.route(
"/__byok-api__/api/settings/tab",
get(settings::get_tab).put(settings::update_tab),
)
.route(
"/__byok-api__/api/settings/desktop",
get(settings::get_desktop).put(settings::update_desktop),
)
.route(
"/__byok-api__/api/harness/cursor/status",
get(harness::status),
)
.route(
"/__byok-api__/api/harness/cursor/ca/initialize",
post(harness::initialize_ca),
)
.route(
"/__byok-api__/api/harness/cursor/enabled",
put(harness::set_enabled),
)
.with_state(service)
.layer(desktop_cors())
}
fn desktop_cors() -> CorsLayer {
CorsLayer::new()
.allow_origin(AllowOrigin::predicate(|origin, _| local_origin(origin)))
.allow_methods([Method::GET, Method::POST, Method::PUT, Method::DELETE])
.allow_headers([
CONTENT_TYPE,
header::ACCEPT_LANGUAGE,
header::HeaderName::from_static("disable-ad-ids"),
])
}
fn local_origin(origin: &HeaderValue) -> bool {
let Ok(origin) = origin.to_str() else {
return false;
};
if origin.eq_ignore_ascii_case("tauri://localhost") {
return true;
}
let Ok(origin) = Url::parse(origin) else {
return false;
};
if !matches!(origin.scheme(), "http" | "https")
|| !origin.username().is_empty()
|| origin.password().is_some()
|| origin.path() != "/"
|| origin.query().is_some()
|| origin.fragment().is_some()
{
return false;
}
match origin.host() {
Some(Host::Domain(host)) => {
host.eq_ignore_ascii_case("localhost") || host.eq_ignore_ascii_case("tauri.localhost")
}
Some(Host::Ipv4(address)) => {
address.is_loopback() || address.is_private() || address.is_link_local()
}
Some(Host::Ipv6(address)) => {
address.is_loopback() || address.is_unique_local() || address.is_unicast_link_local()
}
None => false,
}
}
+98
View File
@@ -0,0 +1,98 @@
//! Implements model configuration endpoints.
use axum::{
extract::{Path, State},
http::StatusCode,
Json,
};
use serde::Deserialize;
use crate::{
model::{ModelConfig, ModelConfigInput},
Result,
};
use super::{
ControlService, DiscoveredModels, LegacyModelImportPreview, LegacyModelImportResult,
ModelConnectivityResult, ModelDiscoveryInput,
};
#[derive(Deserialize)]
pub struct SaveModels {
pub models: Vec<ModelConfigInput>,
}
#[derive(Deserialize)]
pub struct ModelOrder {
pub model_hashes: Vec<String>,
}
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ModelConfig>>> {
Ok(Json(service.models().await?))
}
pub async fn create(
State(service): State<ControlService>,
Json(input): Json<SaveModels>,
) -> Result<(StatusCode, Json<Vec<ModelConfig>>)> {
Ok((
StatusCode::CREATED,
Json(service.create_models(&input.models).await?),
))
}
pub async fn reorder(
State(service): State<ControlService>,
Json(input): Json<ModelOrder>,
) -> Result<Json<Vec<ModelConfig>>> {
Ok(Json(service.reorder_models(&input.model_hashes).await?))
}
pub async fn remove(
State(service): State<ControlService>,
Path(model_hash): Path<String>,
) -> Result<StatusCode> {
service.delete_model(&model_hash).await?;
Ok(StatusCode::NO_CONTENT)
}
pub async fn update(
State(service): State<ControlService>,
Path(model_hash): Path<String>,
Json(input): Json<ModelConfigInput>,
) -> Result<Json<ModelConfig>> {
Ok(Json(service.update_model(&model_hash, &input).await?))
}
pub async fn test(
State(service): State<ControlService>,
Path((model_hash, test_id)): Path<(String, String)>,
) -> Result<Json<ModelConnectivityResult>> {
Ok(Json(service.test_model(&model_hash, &test_id).await?))
}
pub async fn cancel(
State(service): State<ControlService>,
Path((_model_hash, test_id)): Path<(String, String)>,
) -> Result<StatusCode> {
service.cancel_model_test(&test_id);
Ok(StatusCode::NO_CONTENT)
}
pub async fn discover(
State(service): State<ControlService>,
Json(input): Json<ModelDiscoveryInput>,
) -> Result<Json<DiscoveredModels>> {
Ok(Json(service.discover_models(&input).await?))
}
pub async fn import_v0049(
State(service): State<ControlService>,
) -> Result<Json<LegacyModelImportResult>> {
Ok(Json(service.import_v0049_models().await?))
}
pub async fn preview_v0049(
State(service): State<ControlService>,
) -> Result<Json<LegacyModelImportPreview>> {
Ok(Json(service.preview_v0049_models().await?))
}
+30
View File
@@ -0,0 +1,30 @@
//! Implements control dashboard overview endpoints.
//! HTTP handler for the desktop overview aggregates.
use axum::{
extract::{Query, State},
Json,
};
use serde::Deserialize;
use crate::{model::Overview, Result};
use super::ControlService;
#[derive(Debug, Default, Deserialize)]
pub struct OverviewRange {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<String>,
}
pub async fn get(
State(service): State<ControlService>,
Query(range): Query<OverviewRange>,
) -> Result<Json<Overview>> {
Ok(Json(
service
.overview(range.start_ms, range.end_ms, range.model_hashes.as_deref())
.await?,
))
}
+932
View File
@@ -0,0 +1,932 @@
//! Implements control API routing and shared state.
use std::{
collections::{BTreeMap, BTreeSet},
sync::{Arc, Mutex},
time::Instant,
};
use base64::{engine::general_purpose::STANDARD, Engine};
use futures_util::StreamExt;
use reqwest::header::{HeaderName, HeaderValue};
use serde::{Deserialize, Serialize};
use tokio_util::sync::CancellationToken;
use url::Url;
use super::ads::{
AdDismissalInput, AdRuntime, ADS_ENDPOINT, APP_VERSION_HEADER, DEVICE_ID_HEADER,
DISABLED_AD_IDS_HEADER, LANGUAGE_HEADER, OS_HEADER,
};
use crate::{
local_app::CursorHarness,
model::{
ContentPart, CursorRunTraceArtifact, CursorRunTraceSummary, LlmCallRequest,
LlmCallResponseChunk, LlmCallSummary, ModelConfig, ModelConfigInput, ModelInvocation,
ModelRequest, ModelSpec, ModelType, Overview, ProjectedContent, ProjectedMessage,
PromptSpec, ProviderType, Role,
},
provider::{is_valid_response_event, ModelEvent, Provider},
store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store,
TabSettings,
},
Error, Result,
};
#[derive(Clone)]
pub struct ControlService {
store: Store,
cursor_harness: CursorHarness,
provider: Arc<dyn Provider>,
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
}
#[derive(Clone, Debug, Serialize)]
pub struct DiscoveredModels {
pub models: Vec<String>,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportResult {
pub imported: usize,
pub skipped: usize,
pub total: usize,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportPreview {
pub source: String,
pub total: usize,
pub new_models: usize,
pub existing_models: usize,
pub models: Vec<LegacyModelImportPreviewItem>,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportPreviewItem {
pub model_hash: String,
pub display_name: String,
pub model_id: String,
#[serde(rename = "type")]
pub model_type: ModelType,
pub existing: bool,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ModelDiscoveryInput {
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
pub api_key: String,
#[serde(default)]
pub custom_headers_enabled: bool,
#[serde(default = "empty_json_object")]
pub custom_headers: serde_json::Value,
}
fn empty_json_object() -> serde_json::Value {
serde_json::json!({})
}
fn empty_json_object_ref() -> &'static serde_json::Value {
static EMPTY: std::sync::OnceLock<serde_json::Value> = std::sync::OnceLock::new();
EMPTY.get_or_init(empty_json_object)
}
#[derive(Clone, Debug, Serialize)]
pub struct ModelConnectivityResult {
pub duration_ms: u64,
pub first_valid_response_ms: Option<u64>,
pub output_tokens: u64,
pub tokens_per_second: f64,
pub tokens_estimated: bool,
pub output: String,
}
#[derive(Clone, Debug, Serialize)]
pub struct CallDetail {
pub call: CallSummary,
pub request: Option<LlmCallRequest>,
pub response_chunks: Vec<LlmCallResponseChunk>,
pub cursor_trace: Option<CursorTraceDetail>,
}
#[derive(Clone, Debug, Serialize)]
pub struct CallSummary {
#[serde(flatten)]
pub call: LlmCallSummary,
pub call_kind: &'static str,
pub route: &'static str,
}
#[derive(Clone, Debug, Serialize)]
pub struct CursorTraceDetail {
pub trace: CursorRunTraceSummary,
pub artifacts: Vec<CursorTraceArtifactDetail>,
}
#[derive(Clone, Debug, Serialize)]
pub struct CursorTraceArtifactDetail {
pub seq: i64,
pub artifact_type: String,
pub source: String,
pub metadata: serde_json::Value,
pub created_at_ms: i64,
pub byte_count: usize,
pub encoding: &'static str,
pub data: String,
}
#[derive(Clone, Copy, Debug, Deserialize, Serialize)]
pub struct ObservabilitySettings {
pub detailed: bool,
}
impl ControlService {
pub fn new(store: Store, provider: Arc<dyn Provider>) -> Result<Self> {
Ok(Self {
cursor_harness: CursorHarness::new(store.clone())?,
store,
provider,
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
})
}
pub fn cursor_harness(&self) -> &CursorHarness {
&self.cursor_harness
}
pub(super) async fn ads(
&self,
disabled_ad_ids: Option<&str>,
language: &str,
) -> Result<AdRuntime> {
let client = crate::network::client(&self.store).await?;
let installation_id = self.store.installation_id().await?;
let mut request = client
.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(LANGUAGE_HEADER, language)
.timeout(std::time::Duration::from_secs(5));
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
request = request.header(DISABLED_AD_IDS_HEADER, disabled_ad_ids);
}
let response = request.send().await?;
let status = response.status();
if !status.is_success() {
let message = response.text().await.unwrap_or_default();
return Err(Error::Provider(format!(
"advertisement service failed ({status}): {}",
message.chars().take(200).collect::<String>()
)));
}
response.json::<AdRuntime>().await?.into_menu_slots()
}
pub(super) async fn dismiss_ad(&self, ad_id: &str, input: &AdDismissalInput) -> Result<()> {
let client = crate::network::client(&self.store).await?;
let installation_id = self.store.installation_id().await?;
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
Error::Config(format!("advertisement endpoint is invalid: {error}"))
})?;
endpoint.set_query(None);
endpoint
.path_segments_mut()
.map_err(|_| Error::Config("advertisement endpoint cannot contain an ad id".into()))?
.push(ad_id)
.push("dismissals");
let response = client
.post(endpoint)
.header(DEVICE_ID_HEADER, installation_id)
.header(OS_HEADER, std::env::consts::OS)
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
.json(input)
.timeout(std::time::Duration::from_secs(5))
.send()
.await?;
let status = response.status();
if !status.is_success() {
let message = response.text().await.unwrap_or_default();
return Err(Error::Provider(format!(
"advertisement dismissal failed ({status}): {}",
message.chars().take(200).collect::<String>()
)));
}
Ok(())
}
pub async fn models(&self) -> Result<Vec<ModelConfig>> {
self.store.models().await
}
pub async fn overview(
&self,
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<&str>,
) -> Result<Overview> {
self.store.overview(start_ms, end_ms, model_hashes).await
}
pub async fn create_models(&self, models: &[ModelConfigInput]) -> Result<Vec<ModelConfig>> {
self.store.create_models(models).await
}
pub async fn reorder_models(&self, model_hashes: &[String]) -> Result<Vec<ModelConfig>> {
self.store.reorder_models(model_hashes).await
}
pub async fn delete_model(&self, model_hash: &str) -> Result<()> {
self.store.delete_model(model_hash).await
}
pub async fn update_model(
&self,
model_hash: &str,
input: &ModelConfigInput,
) -> Result<ModelConfig> {
self.store.update_model(model_hash, input).await
}
pub async fn test_model(
&self,
model_hash: &str,
test_id: &str,
) -> Result<ModelConnectivityResult> {
let cancellation = CancellationToken::new();
let cancellation = {
let mut tests = self
.model_tests
.lock()
.expect("model test registry mutex poisoned");
tests
.entry(test_id.to_owned())
.or_insert_with(|| cancellation.clone())
.clone()
};
let result = self.run_model_test(model_hash, cancellation).await;
self.model_tests
.lock()
.expect("model test registry mutex poisoned")
.remove(test_id);
result
}
pub fn cancel_model_test(&self, test_id: &str) {
let cancellation = {
let mut tests = self
.model_tests
.lock()
.expect("model test registry mutex poisoned");
tests.entry(test_id.to_owned()).or_default().clone()
};
cancellation.cancel();
}
async fn run_model_test(
&self,
model_hash: &str,
cancellation: CancellationToken,
) -> Result<ModelConnectivityResult> {
const TEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(45);
const TEST_PROMPT: &str = "Output the numbers 1 through 120 separated by a single space. No commas, no newlines, no explanation.";
let configured = self
.store
.model(model_hash)
.await?
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
let mut model = ModelSpec::new(model_hash);
configured.configure(&mut model);
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
let call_id = format!("model-test-{}", uuid::Uuid::new_v4());
let invocation = ModelInvocation {
call_id: call_id.clone(),
run_id: call_id.clone(),
conversation_id: call_id.clone(),
provider_call_index: 0,
request: ModelRequest {
prompt: PromptSpec {
instructions: String::new(),
tools: Vec::new(),
},
model,
history: vec![ProjectedMessage {
message_id: "connectivity-test".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text {
text: TEST_PROMPT.into(),
}]),
}],
},
};
let started = Instant::now();
let mut first_valid_response_at = None;
let mut output_tokens = None;
let mut output = String::new();
let stream = self.provider.stream(invocation, cancellation.clone());
let completed = tokio::time::timeout(TEST_TIMEOUT, async {
futures_util::pin_mut!(stream);
let mut finished = false;
while let Some(event) = stream.next().await {
let event = event?;
if first_valid_response_at.is_none() && is_valid_response_event(&event) {
first_valid_response_at = Some(Instant::now());
}
match event {
ModelEvent::TextDelta(delta) => {
output.push_str(&delta);
}
ModelEvent::Usage(usage) => {
if let Some(tokens) = usage.output_tokens.filter(|tokens| *tokens > 0) {
output_tokens = Some(
output_tokens.map_or(tokens, |current: u64| current.max(tokens)),
);
}
}
ModelEvent::Done(_) => finished = true,
_ => {}
}
}
if cancellation.is_cancelled() {
return Err(Error::Cancelled);
}
if !finished {
return Err(Error::Protocol(
"provider stream ended without Done during connectivity test".into(),
));
}
Ok(())
})
.await;
match completed {
Ok(result) => result?,
Err(_) => {
cancellation.cancel();
self.store
.finish_llm_call(
&call_id,
"error",
None,
started.elapsed().as_millis().min(i64::MAX as u128) as i64,
Some("timeout"),
Some("model connectivity test timed out after 45 seconds"),
)
.await?;
return Err(Error::Provider(
"model connectivity test timed out after 45 seconds".into(),
));
}
}
let elapsed = started.elapsed();
let output = output.trim().to_string();
if first_valid_response_at.is_none() {
return Err(Error::Provider(
"model connectivity test received no valid response".into(),
));
}
let tokens_estimated = output_tokens.is_none();
let output_tokens = output_tokens.unwrap_or_else(|| estimate_output_tokens(&output));
Ok(ModelConnectivityResult {
duration_ms: elapsed.as_millis().min(u128::from(u64::MAX)) as u64,
first_valid_response_ms: first_valid_response_at.map(|first| {
first
.duration_since(started)
.as_millis()
.min(u128::from(u64::MAX)) as u64
}),
output_tokens,
tokens_per_second: if elapsed.is_zero() {
0.0
} else {
output_tokens as f64 / elapsed.as_secs_f64()
},
tokens_estimated,
output,
})
}
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let base_url = crate::model::normalize_request_url(&input.base_url)?;
discover_models_from_endpoint(
&client,
match input.model_type {
ModelType::OpenAi => ProviderType::OpenAiResponses,
ModelType::Anthropic => ProviderType::Anthropic,
},
&base_url,
&input.api_key,
if input.custom_headers_enabled {
&input.custom_headers
} else {
empty_json_object_ref()
},
)
.await
}
pub async fn import_v0049_models(&self) -> Result<LegacyModelImportResult> {
let path = crate::config::v0049_config_path()?;
let outcome = self.store.import_v0049_model_config(&path).await?;
Ok(LegacyModelImportResult {
imported: outcome.imported,
skipped: outcome.skipped,
total: outcome.total,
})
}
pub async fn preview_v0049_models(&self) -> Result<LegacyModelImportPreview> {
let path = crate::config::v0049_config_path()?;
let plan = self.store.preview_v0049_model_config(&path).await?;
let total = plan.models.len();
let existing_models = plan.models.iter().filter(|model| model.existing).count();
Ok(LegacyModelImportPreview {
source: path.display().to_string(),
total,
new_models: total - existing_models,
existing_models,
models: plan
.models
.into_iter()
.map(|model| LegacyModelImportPreviewItem {
model_hash: model.model_hash,
display_name: model.input.display_name,
model_id: model.input.model_id,
model_type: model.input.model_type,
existing: model.existing,
})
.collect(),
})
}
pub async fn calls(&self, limit: i64) -> Result<Vec<CallSummary>> {
let mut calls = self
.store
.llm_calls(limit)
.await?
.into_iter()
.map(|call| CallSummary {
call,
call_kind: "provider_llm",
route: "local_byok",
})
.collect::<Vec<_>>();
calls.extend(
self.store
.official_cursor_traces(limit)
.await?
.into_iter()
.map(official_call),
);
calls.sort_by_key(|call| std::cmp::Reverse(call.call.created_at_ms));
calls.truncate(limit.clamp(1, 500) as usize);
Ok(calls)
}
pub async fn call(&self, call_id: &str) -> Result<CallDetail> {
if let Some(call) = self.store.llm_call(call_id).await? {
let cursor_trace = self.cursor_trace_detail(&call.run_id).await?;
return Ok(CallDetail {
request: self.store.llm_call_request(call_id).await?,
response_chunks: self.store.llm_call_chunks(call_id).await?,
call: CallSummary {
call,
call_kind: "provider_llm",
route: "local_byok",
},
cursor_trace,
});
}
let request_id = call_id.strip_prefix("cursor:").unwrap_or(call_id);
let trace = self
.store
.cursor_trace(request_id)
.await?
.filter(|trace| trace.route == "cursor_official")
.ok_or_else(|| Error::RunNotFound(format!("call {call_id}")))?;
Ok(CallDetail {
call: official_call(trace.clone()),
request: None,
response_chunks: Vec::new(),
cursor_trace: Some(self.cursor_trace_detail_from(trace).await?),
})
}
async fn cursor_trace_detail(&self, request_id: &str) -> Result<Option<CursorTraceDetail>> {
let Some(trace) = self.store.cursor_trace(request_id).await? else {
return Ok(None);
};
Ok(Some(self.cursor_trace_detail_from(trace).await?))
}
async fn cursor_trace_detail_from(
&self,
trace: CursorRunTraceSummary,
) -> Result<CursorTraceDetail> {
let artifacts = self
.store
.cursor_trace_artifacts(&trace.request_id)
.await?
.into_iter()
.map(cursor_artifact)
.collect();
Ok(CursorTraceDetail { trace, artifacts })
}
pub async fn observability(&self) -> Result<ObservabilitySettings> {
Ok(ObservabilitySettings {
detailed: self.store.detailed_logging().await?,
})
}
pub async fn set_observability(
&self,
settings: ObservabilitySettings,
) -> Result<ObservabilitySettings> {
self.store.set_detailed_logging(settings.detailed).await?;
Ok(settings)
}
pub async fn ports(&self) -> Result<PortSettings> {
self.store.port_settings().await
}
pub async fn set_ports(&self, settings: PortSettings) -> Result<PortSettings> {
self.store.set_port_settings(settings).await?;
Ok(settings)
}
pub async fn statistics_storage(&self) -> Result<StatisticsStorage> {
self.store.statistics_storage().await
}
pub async fn clear_statistics_storage(&self) -> Result<StatisticsStorage> {
self.store.clear_statistics_storage().await
}
pub async fn clear_all_statistics_storage(&self) -> Result<StatisticsStorage> {
self.store.clear_all_statistics_storage().await
}
pub async fn proxy_settings(&self) -> Result<ProxySettings> {
self.store.proxy_settings().await
}
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
self.store.set_proxy_settings(settings).await
}
pub async fn tab_settings(&self) -> Result<TabSettings> {
self.store.tab_settings().await
}
pub async fn set_tab_settings(&self, settings: TabSettings) -> Result<TabSettings> {
self.cursor_harness.set_tab_settings(settings).await
}
pub async fn desktop_settings(&self) -> Result<DesktopSettings> {
self.store.desktop_settings().await
}
pub async fn set_desktop_settings(&self, settings: DesktopSettings) -> Result<()> {
self.store.set_desktop_settings(settings).await
}
}
fn official_call(trace: CursorRunTraceSummary) -> CallSummary {
let model_id = trace.model_id.clone().unwrap_or_else(|| "Cursor".into());
let ttfb = trace
.first_response_at_ms
.map(|value| (value - trace.received_at_ms).max(0));
let duration = trace
.finished_at_ms
.map(|value| (value - trace.received_at_ms).max(0));
let error = trace.error_message.clone();
CallSummary {
call: LlmCallSummary {
call_id: format!("cursor:{}", trace.request_id),
run_id: trace.request_id.clone(),
conversation_id: trace
.conversation_id
.clone()
.unwrap_or_else(|| trace.request_id.clone()),
provider_call_index: 0,
model_hash: None,
provider_type: "cursor-official".into(),
provider_url: "https://api2.cursor.sh".into(),
request_type: "cursor-run-sse".into(),
request_url: "https://api2.cursor.sh/agent.v1.AgentService/RunSSE".into(),
model_id: model_id.clone(),
display_name: model_id,
reasoning_effort: None,
fast: None,
status: trace.status.clone(),
finish_reason: None,
created_at_ms: trace.received_at_ms,
request_started_at_ms: Some(trace.received_at_ms),
response_headers_at_ms: trace.first_response_at_ms,
first_event_at_ms: trace.first_response_at_ms,
first_text_at_ms: None,
first_valid_response_at_ms: None,
finished_at_ms: trace.finished_at_ms,
queue_ms: None,
ttfb_ms: ttfb,
ttft_ms: None,
ttfr_ms: None,
duration_ms: duration,
input_tokens: None,
output_tokens: None,
total_tokens: None,
cache_read_tokens: None,
cache_write_tokens: None,
reasoning_tokens: None,
usage: None,
message_count: 0,
tool_count: 0,
request_bytes: Some(trace.request_bytes),
response_bytes: trace.response_bytes,
stream_event_count: trace.response_event_count,
http_status: trace.http_status,
error_kind: error.as_ref().map(|_| "cursor_official".into()),
error_message: error,
detailed: true,
},
call_kind: "cursor_official",
route: "cursor_official",
}
}
fn cursor_artifact(artifact: CursorRunTraceArtifact) -> CursorTraceArtifactDetail {
let byte_count = artifact.data.len();
let (encoding, data) = match readable_utf8(&artifact.data) {
Some(value) => ("utf8", value.into()),
None => ("base64", STANDARD.encode(&artifact.data)),
};
CursorTraceArtifactDetail {
seq: artifact.seq,
artifact_type: artifact.artifact_type,
source: artifact.source,
metadata: artifact.metadata,
created_at_ms: artifact.created_at_ms,
byte_count,
encoding,
data,
}
}
fn readable_utf8(data: &[u8]) -> Option<&str> {
let value = std::str::from_utf8(data).ok()?;
value
.chars()
.all(|character| !character.is_control() || matches!(character, '\n' | '\r' | '\t'))
.then_some(value)
}
async fn discover_models_from_endpoint(
client: &reqwest::Client,
provider_type: ProviderType,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<DiscoveredModels> {
let mut models = match provider_type {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
openai_models(client, base_url, api_key, custom_headers).await?
}
ProviderType::Anthropic => {
anthropic_models(client, base_url, api_key, custom_headers).await?
}
};
models.sort();
models.dedup();
Ok(DiscoveredModels { models })
}
fn model_discovery_url(base_url: &str) -> Result<Url> {
let mut url = Url::parse(base_url)
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
if url.host_str().is_none() {
return Err(Error::Config(
"model request URL must contain a host".into(),
));
}
// 在现有路径上追加,而不是整段替换:多数编程套餐的 API 挂在子路径下
// (/api/anthropic、/coding、/api/paas/v4 等),直接 set_path("/v1/models")
// 会把这些前缀吃掉,发现请求必然 404
let path = url.path().trim_end_matches('/');
let last = path.rsplit('/').next().unwrap_or("");
let versioned = last.len() > 1
&& last.starts_with('v')
&& last[1..].bytes().all(|byte| byte.is_ascii_digit());
let new_path = if let Some(parent) = path.strip_suffix("/chat/completions") {
// 完整请求 URL:剥掉端点段(chat/completions 是两段),换成 models
format!("{parent}/models")
} else if let Some(parent) = path
.strip_suffix("/responses")
.or_else(|| path.strip_suffix("/messages"))
.or_else(|| path.strip_suffix("/completions"))
{
format!("{parent}/models")
} else if path.is_empty() {
"/v1/models".to_string()
} else if versioned {
// 已带版本段(/v1、/api/v3、/api/paas/v4):只补 models
format!("{path}/models")
} else {
format!("{path}/v1/models")
};
url.set_path(&new_path);
url.set_query(None);
url.set_fragment(None);
Ok(url)
}
fn model_discovery_urls(base_url: &str) -> Result<Vec<Url>> {
let mut configured = Url::parse(base_url)
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
let path = configured.path().trim_end_matches('/');
let tail = path.rsplit('/').next().unwrap_or_default();
if matches!(tail.to_ascii_lowercase().as_str(), "model" | "models") {
configured.set_query(None);
configured.set_fragment(None);
return Ok(vec![configured]);
}
let primary = model_discovery_url(base_url)?;
let versioned = tail.len() > 1
&& tail.starts_with('v')
&& tail[1..].bytes().all(|byte| byte.is_ascii_digit());
let complete_request_url = [
"/chat/completions",
"/responses",
"/messages",
"/completions",
]
.iter()
.any(|suffix| path.to_ascii_lowercase().ends_with(suffix));
if versioned || complete_request_url {
return Ok(vec![primary]);
}
let Some(prefix) = primary.path().strip_suffix("/v1/models") else {
return Ok(vec![primary]);
};
let mut fallback = primary.clone();
fallback.set_path(&format!("{prefix}/models"));
Ok(vec![primary, fallback])
}
async fn openai_models(
client: &reqwest::Client,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut last_error = None;
for url in model_discovery_urls(base_url)? {
match openai_models_at(client, url, api_key, custom_headers).await {
Ok(models) => return Ok(models),
Err(error) => last_error = Some(error),
}
}
Err(last_error.unwrap_or_else(|| Error::Provider("no model discovery URL available".into())))
}
async fn openai_models_at(
client: &reqwest::Client,
url: Url,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut request = client.get(url);
if !api_key.is_empty() {
request = request.bearer_auth(api_key);
}
let response = apply_discovery_headers(request, custom_headers)?
.send()
.await?;
let status = response.status();
let body: serde_json::Value = response.json().await?;
if !status.is_success() {
return Err(Error::Provider(format!(
"model discovery failed ({status}): {body}"
)));
}
Ok(model_ids(body.get("data").unwrap_or(&body)))
}
async fn anthropic_models(
client: &reqwest::Client,
base_url: &str,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut last_error = None;
for url in model_discovery_urls(base_url)? {
match anthropic_models_at(client, url, api_key, custom_headers).await {
Ok(models) => return Ok(models),
Err(error) => last_error = Some(error),
}
}
Err(last_error.unwrap_or_else(|| Error::Provider("no model discovery URL available".into())))
}
async fn anthropic_models_at(
client: &reqwest::Client,
url: Url,
api_key: &str,
custom_headers: &serde_json::Value,
) -> Result<Vec<String>> {
let mut after_id = None::<String>;
let mut found = BTreeSet::new();
loop {
let mut request = client
.get(url.clone())
.query(&[("limit", "100")])
.header("anthropic-version", "2023-06-01");
if !api_key.is_empty() {
request = request.header("x-api-key", api_key);
}
if let Some(after_id) = &after_id {
request = request.query(&[("after_id", after_id)]);
}
let response = apply_discovery_headers(request, custom_headers)?
.send()
.await?;
let status = response.status();
let body: serde_json::Value = response.json().await?;
if !status.is_success() {
return Err(Error::Provider(format!(
"model discovery failed ({status}): {body}"
)));
}
found.extend(model_ids(body.get("data").unwrap_or(&body)));
if body.get("has_more").and_then(serde_json::Value::as_bool) != Some(true) {
break;
}
after_id = body
.get("last_id")
.and_then(serde_json::Value::as_str)
.map(str::to_owned);
if after_id.is_none() {
return Err(Error::Provider(
"Anthropic model response has_more without last_id".into(),
));
}
}
Ok(found.into_iter().collect())
}
fn model_ids(value: &serde_json::Value) -> Vec<String> {
value
.as_array()
.into_iter()
.flatten()
.filter_map(|item| match item {
serde_json::Value::String(id) => Some(id.clone()),
serde_json::Value::Object(object) => object
.get("id")
.or_else(|| object.get("name"))
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
_ => None,
})
.collect()
}
fn estimate_output_tokens(output: &str) -> u64 {
let words = output.split_whitespace().count() as u64;
if words > 0 {
words
} else if output.is_empty() {
0
} else {
(output.chars().count() as u64).div_ceil(4)
}
}
fn apply_discovery_headers(
mut request: reqwest::RequestBuilder,
headers: &serde_json::Value,
) -> Result<reqwest::RequestBuilder> {
let object = headers
.as_object()
.ok_or_else(|| Error::Config("custom headers must be an object".into()))?;
for (name, value) in object {
if name.eq_ignore_ascii_case("user-agent") {
continue;
}
let value = value
.as_str()
.ok_or_else(|| Error::Config(format!("custom header {name} must be a string")))?;
let name = HeaderName::try_from(name)
.map_err(|error| Error::Config(format!("invalid header name: {error}")))?;
let value = HeaderValue::try_from(value)
.map_err(|error| Error::Config(format!("invalid header value: {error}")))?;
request = request.header(name, value);
}
Ok(request)
}
+89
View File
@@ -0,0 +1,89 @@
//! Implements settings management endpoints.
use crate::Result;
use axum::{extract::State, Json};
use serde::Deserialize;
use crate::store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage,
StatisticsStorageScope, TabSettings,
};
use super::{ControlService, ObservabilitySettings};
pub async fn get(State(service): State<ControlService>) -> Result<Json<ObservabilitySettings>> {
Ok(Json(service.observability().await?))
}
pub async fn update(
State(service): State<ControlService>,
Json(settings): Json<ObservabilitySettings>,
) -> Result<Json<ObservabilitySettings>> {
Ok(Json(service.set_observability(settings).await?))
}
pub async fn get_ports(State(service): State<ControlService>) -> Result<Json<PortSettings>> {
Ok(Json(service.ports().await?))
}
pub async fn update_ports(
State(service): State<ControlService>,
Json(settings): Json<PortSettings>,
) -> Result<Json<PortSettings>> {
Ok(Json(service.set_ports(settings).await?))
}
pub async fn get_storage(State(service): State<ControlService>) -> Result<Json<StatisticsStorage>> {
Ok(Json(service.statistics_storage().await?))
}
pub async fn clear_storage(
State(service): State<ControlService>,
input: Option<Json<ClearStorageInput>>,
) -> Result<Json<StatisticsStorage>> {
let scope = input.map(|Json(input)| input.scope).unwrap_or_default();
let storage = match scope {
StatisticsStorageScope::Details => service.clear_statistics_storage().await?,
StatisticsStorageScope::All => service.clear_all_statistics_storage().await?,
};
Ok(Json(storage))
}
#[derive(Deserialize)]
pub struct ClearStorageInput {
#[serde(default)]
pub scope: StatisticsStorageScope,
}
pub async fn get_proxy(State(service): State<ControlService>) -> Result<Json<ProxySettings>> {
Ok(Json(service.proxy_settings().await?))
}
pub async fn update_proxy(
State(service): State<ControlService>,
Json(settings): Json<ProxySettingsInput>,
) -> Result<Json<ProxySettings>> {
Ok(Json(service.set_proxy_settings(settings).await?))
}
pub async fn get_tab(State(service): State<ControlService>) -> Result<Json<TabSettings>> {
Ok(Json(service.tab_settings().await?))
}
pub async fn update_tab(
State(service): State<ControlService>,
Json(settings): Json<TabSettings>,
) -> Result<Json<TabSettings>> {
Ok(Json(service.set_tab_settings(settings).await?))
}
pub async fn get_desktop(State(service): State<ControlService>) -> Result<Json<DesktopSettings>> {
Ok(Json(service.desktop_settings().await?))
}
pub async fn update_desktop(
State(service): State<ControlService>,
Json(settings): Json<DesktopSettings>,
) -> Result<Json<DesktopSettings>> {
service.set_desktop_settings(settings).await?;
get_desktop(State(service)).await
}
+301
View File
@@ -0,0 +1,301 @@
//! Coordinates construction of a complete Cursor checkpoint.
use std::collections::HashSet;
use prost::Message;
use crate::{
cursor::{
checkpoint::{messages, PendingSteps},
protocol::proto::agent::v1 as pb,
services::blob_sync::BlobSynchronizer,
transport::TransportHandle,
},
model::{CanonicalMessage, ToolCall, ToolDefinition, ToolRoundAssistant},
store::Store,
Result,
};
use super::{derived, roots::RootFrontier, turns::TurnFrontier};
#[derive(Clone)]
pub struct CheckpointBuilder {
pub(super) store: Store,
pub(super) sync: BlobSynchronizer,
pub(super) parent_tool_call_id: Option<String>,
pub(super) base: pb::ConversationStateStructure,
pub(super) model: String,
pub(super) max_context_tokens: Option<u64>,
pub(super) instructions: String,
pub(super) tool_definitions: Vec<ToolDefinition>,
pub(super) allowed_tools: Vec<String>,
pub(super) dynamic_tools: HashSet<String>,
pub(super) turn_user: Option<pb::UserMessage>,
pub(super) roots: Option<RootFrontier>,
pub(super) turn: Option<TurnFrontier>,
pub(super) turns_initialized: bool,
}
impl CheckpointBuilder {
pub fn new(
store: Store,
sync: BlobSynchronizer,
parent_tool_call_id: Option<String>,
base: Option<pb::ConversationStateStructure>,
) -> Self {
Self {
store,
sync,
parent_tool_call_id,
base: base.unwrap_or_default(),
model: String::new(),
max_context_tokens: None,
instructions: String::new(),
tool_definitions: Vec::new(),
allowed_tools: Vec::new(),
dynamic_tools: HashSet::new(),
turn_user: None,
roots: None,
turn: None,
turns_initialized: false,
}
}
pub fn configure(
&mut self,
model: String,
max_context_tokens: Option<u64>,
instructions: String,
tool_definitions: Vec<ToolDefinition>,
dynamic_tools: HashSet<String>,
turn_user: Option<pb::UserMessage>,
) {
self.model = model;
self.max_context_tokens = max_context_tokens;
self.instructions = instructions;
self.allowed_tools = tool_definitions
.iter()
.map(|tool| tool.name.clone())
.collect();
self.tool_definitions = tool_definitions;
self.dynamic_tools = dynamic_tools;
self.turn_user = turn_user;
}
pub(crate) fn record_context_tokens(&mut self, used_tokens: Option<u64>) {
let previous = self
.base
.token_details
.as_ref()
.map(|details| details.max_tokens as u64);
let max_tokens = context_limit(self.max_context_tokens, previous);
let Some(max_tokens) = max_tokens else {
return;
};
let details = self.base.token_details.get_or_insert_with(Default::default);
if let Some(used_tokens) = used_tokens {
details.used_tokens = used_tokens.min(u32::MAX as u64) as u32;
}
details.max_tokens = max_tokens.min(u32::MAX as u64) as u32;
details.prompt_context_usage_tree = None;
details.prompt_context_usage_snapshot_blob_id = None;
}
pub async fn settled(
&mut self,
messages: &[CanonicalMessage],
mode: i32,
presentation: &PendingSteps,
) -> Result<pb::ConversationStateStructure> {
self.build_state(messages, mode, Vec::new(), presentation)
.await
}
pub async fn staged_tool_round(
&mut self,
stable_messages: &[CanonicalMessage],
mode: i32,
assistant: &ToolRoundAssistant,
calls: &[ToolCall],
started_at_ms: u64,
presentation: &PendingSteps,
) -> Result<pb::ConversationStateStructure> {
let pending = messages::staged_tool_round(
assistant,
calls,
&self.model,
&self.allowed_tools,
&self.dynamic_tools,
started_at_ms,
)?;
self.build_state(stable_messages, mode, vec![pending], presentation)
.await
}
pub async fn staged_final(
&mut self,
stable_messages: &[CanonicalMessage],
mode: i32,
assistant: &CanonicalMessage,
started_at_ms: u64,
presentation: &PendingSteps,
) -> Result<pb::ConversationStateStructure> {
let pending = messages::staged_final(
assistant,
&self.model,
&self.allowed_tools,
&self.dynamic_tools,
started_at_ms,
)?;
self.build_state(stable_messages, mode, vec![pending], presentation)
.await
}
async fn build_state(
&mut self,
messages: &[CanonicalMessage],
mode: i32,
pending_tool_calls: Vec<String>,
presentation: &PendingSteps,
) -> Result<pb::ConversationStateStructure> {
self.record_background_subagents(presentation);
let root_ids = self.project_roots(messages).await?;
let turn_ids = self.project_turns(mode, presentation).await?;
let (todo_ids, plan_id) = self.build_derived_state(messages).await?;
self.base.todos = todo_ids.iter().map(|id| id.as_bytes().to_vec()).collect();
self.base.plan = plan_id.as_ref().map(|id| id.as_bytes().to_vec());
let communicate_update_states_by_parent_tool_call_id = self
.parent_tool_call_id
.as_ref()
.and_then(|parent| {
derived::update_current_step_state(messages).map(|state| (parent.clone(), state))
})
.into_iter()
.collect();
for path in &presentation.read_paths {
if !self.base.read_paths.contains(path) {
self.base.read_paths.push(path.clone());
}
}
let mut checkpoint = self.base.clone();
checkpoint.root_prompt_messages_json =
root_ids.iter().map(|id| id.as_bytes().to_vec()).collect();
checkpoint.turns = turn_ids.iter().map(|id| id.as_bytes().to_vec()).collect();
checkpoint.pending_tool_calls = pending_tool_calls;
checkpoint.mode = Some(mode);
checkpoint.communicate_update_states_by_parent_tool_call_id =
communicate_update_states_by_parent_tool_call_id;
if let Some(details) = checkpoint.token_details.as_mut() {
details.breakdown = Some(crate::cursor::services::usage::breakdown(
details.used_tokens,
details.max_tokens,
details.breakdown.as_ref(),
&self.instructions,
&self.tool_definitions,
&self.dynamic_tools,
messages,
)?);
}
Ok(checkpoint)
}
fn record_background_subagents(&mut self, presentation: &PendingSteps) {
for step in &presentation.steps {
let Some(pb::conversation_step::Message::ToolCall(call)) = step.message.as_ref() else {
continue;
};
let Some(pb::tool_call::Tool::TaskToolCall(task)) = call.tool.as_ref() else {
continue;
};
let (Some(args), Some(result)) = (task.args.as_ref(), task.result.as_ref()) else {
continue;
};
let Some(pb::task_result::Result::Success(success)) = result.result.as_ref() else {
continue;
};
if !success.is_background {
continue;
}
let Some(agent_id) = success.agent_id.as_ref().filter(|id| !id.is_empty()) else {
continue;
};
let Some(tool_call_id) = call.tool_call_id.as_ref().filter(|id| !id.is_empty()) else {
continue;
};
let started_at_ms = call
.started_at_ms
.unwrap_or_else(crate::cursor::tools::runtime::now_ms);
let last_used_timestamp_ms = call.completed_at_ms.unwrap_or(started_at_ms);
self.base
.subagent_states
.entry(agent_id.clone())
.and_modify(|state| state.last_used_timestamp_ms = last_used_timestamp_ms)
.or_insert_with(|| pb::SubagentPersistedState {
conversation_state: None,
created_timestamp_ms: started_at_ms,
last_used_timestamp_ms,
subagent_type: args.subagent_type.clone(),
model_id: args.model.clone(),
environment: args.environment,
cloud_subagent: None,
first_class_bc_id: None,
cloud_requested_environment_build_id: None,
machine: args.machine.clone(),
});
self.base.subagent_runs_by_parent_tool_call_id.insert(
tool_call_id.clone(),
pb::SubagentRunState {
parent_tool_call_id: tool_call_id.clone(),
subagent_id: Some(agent_id.clone()),
environment: args.environment,
status: pb::SubagentRunStatus::Backgrounded as i32,
title: Some(args.description.clone()),
detail: success.result_suffix.clone(),
transcript_path: success.transcript_path.clone(),
output_path: None,
completed_timestamp_ms: None,
completion_reason: None,
},
);
}
}
pub async fn publish(
&self,
handle: &TransportHandle,
checkpoint: &pb::ConversationStateStructure,
) -> Result<()> {
tracing::debug!(
request_id = self.sync.request_id(),
stable_roots = checkpoint.root_prompt_messages_json.len(),
pending_assistants = checkpoint.pending_tool_calls.len(),
"publishing Cursor checkpoint"
);
let result = handle.emit(&pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(
pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoint.clone()),
),
});
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;
}
result
}
}
fn context_limit(selected: Option<u64>, previous: Option<u64>) -> Option<u64> {
selected.or(previous.filter(|tokens| *tokens != 0))
}
+170
View File
@@ -0,0 +1,170 @@
//! Derives Todo, Plan, and related checkpoint state from Messages.
use std::collections::HashMap;
use prost::Message;
use crate::{
cursor::{prompting::fold_derived_state, protocol::proto::agent::v1 as pb},
model::{CanonicalMessage, MessageContent},
store::BlobId,
Error, Result,
};
use super::CheckpointBuilder;
impl CheckpointBuilder {
pub(super) async fn build_derived_state(
&self,
messages: &[CanonicalMessage],
) -> Result<(Vec<BlobId>, Option<BlobId>)> {
let state = fold_derived_state(messages);
let todo_values = state
.todos
.as_ref()
.map(|value| {
value
.get("todos")
.and_then(serde_json::Value::as_array)
.ok_or_else(|| Error::Protocol("TodoWrite state is missing todos[]".into()))
})
.transpose()?;
let mut todo_ids = Vec::new();
for (index, todo) in todo_values.into_iter().flatten().enumerate() {
let status = match todo
.get("status")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol("TodoWrite item is missing status".into()))?
{
"in_progress" => pb::TodoStatus::InProgress,
"completed" => pb::TodoStatus::Completed,
"cancelled" => pb::TodoStatus::Cancelled,
"pending" => pb::TodoStatus::Pending,
status => {
return Err(Error::Protocol(format!(
"unknown TodoWrite status: {status}"
)))
}
};
let message = pb::TodoItem {
id: todo
.get("id")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol("TodoWrite item is missing id".into()))?
.into(),
content: todo
.get("content")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol("TodoWrite item is missing content".into()))?
.into(),
status: status as i32,
created_at: 0,
updated_at: 0,
dependencies: todo
.get("dependencies")
.and_then(serde_json::Value::as_array)
.into_iter()
.flatten()
.filter_map(serde_json::Value::as_str)
.map(str::to_string)
.collect(),
};
let mut encoded = Vec::new();
message.encode(&mut encoded)?;
let id = BlobId::digest(&encoded);
if self.base.todos.get(index).map(|raw| raw.as_slice()) == Some(id.as_bytes()) {
todo_ids.push(id);
} else {
todo_ids.push(self.sync.persist(&encoded, &[]).await?);
}
}
let plan_id = if let Some(value) = state.plan {
let text = value
.get("plan")
.and_then(serde_json::Value::as_str)
.or_else(|| value.as_str())
.or_else(|| value.get("overview").and_then(serde_json::Value::as_str))
.ok_or_else(|| Error::Protocol("plan state has no textual plan".into()))?;
let mut encoded = Vec::new();
pb::ConversationPlan { plan: text.into() }.encode(&mut encoded)?;
let id = BlobId::digest(&encoded);
if self.base.plan.as_deref() == Some(id.as_bytes()) {
Some(id)
} else {
Some(self.sync.persist(&encoded, &[]).await?)
}
} else {
None
};
Ok((todo_ids, plan_id))
}
}
pub(super) fn update_current_step_state(
messages: &[CanonicalMessage],
) -> Option<pb::CommunicateUpdateTurnState> {
let result_indices = messages
.iter()
.filter_map(|message| match &message.content {
MessageContent::ToolResult(result) => {
update_message_index(&result.content).map(|index| (result.call_id.as_str(), index))
}
_ => None,
})
.collect::<HashMap<_, _>>();
let mut state = pb::CommunicateUpdateTurnState::default();
for message in messages {
let MessageContent::Assistant { tool_calls, .. } = &message.content else {
continue;
};
for call in tool_calls {
if normalize(&call.name) != "updatecurrentstep" {
continue;
}
if let (Some(step), Some(message_index)) = (
call.arguments
.get("current_step")
.and_then(serde_json::Value::as_str),
result_indices.get(call.call_id.as_str()),
) {
state.history.push(pb::CommunicateUpdateHistoryEntry {
step: step.into(),
message_index: *message_index,
});
}
if let Some(summary) = call
.arguments
.get("final_summary")
.and_then(serde_json::Value::as_str)
{
state.final_summary = Some(summary.into());
}
if let Some(subtitle) = call
.arguments
.get("completed_subtitle")
.and_then(serde_json::Value::as_str)
{
state.completed_subtitle = Some(subtitle.into());
}
}
}
(!state.history.is_empty()
|| state.final_summary.is_some()
|| state.completed_subtitle.is_some())
.then_some(state)
}
fn update_message_index(output: &str) -> Option<u32> {
let value: serde_json::Value = serde_json::from_str(output).ok()?;
value
.get("success")
.and_then(|success| success.get("message_index"))
.and_then(serde_json::Value::as_u64)
.and_then(|index| u32::try_from(index).ok())
}
fn normalize(name: &str) -> String {
name.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
@@ -0,0 +1,252 @@
//! Decodes Cursor checkpoint message data into canonical Messages.
use base64::{engine::general_purpose::STANDARD, Engine};
use serde_json::Value;
use crate::{
model::{
CanonicalMessage, ContentPart, MessageContent, Origin, RecoveredToolRound, Role, ToolCall,
ToolCallContent, ToolResultContent, ToolRoundAssistant, ToolRoundId,
},
store::BlobId,
Error, Result,
};
use super::REPLAY_ENVELOPE_PREFIX;
pub fn decode(data: &[u8], internal_id: String) -> Result<CanonicalMessage> {
let value: Value = serde_json::from_slice(data)?;
let role = match required_string(&value, "role")? {
"system" => Role::System,
"user" => Role::User,
"assistant" => Role::Assistant,
"tool" => Role::Tool,
role => {
return Err(Error::Protocol(format!(
"unknown Cursor message role: {role}"
)))
}
};
let wire_id = value
.get("id")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string();
let is_request_context = role == Role::User && wire_id.starts_with("request-context:");
let is_prompt_context =
is_request_context || role == Role::User && wire_id.starts_with("selected-context:");
let origin = match role {
Role::System => Origin::Prompt,
Role::Assistant => Origin::Assistant,
Role::Tool => Origin::Tool,
Role::User if wire_id.starts_with("runtime:") => Origin::Runtime,
Role::User if is_prompt_context => Origin::Prompt,
Role::User => Origin::User,
};
let runtime_event_id = wire_id.strip_prefix("runtime:").map(str::to_string);
let content = match role {
Role::Assistant => decode_assistant(&value, &internal_id)?,
Role::Tool => MessageContent::ToolResult(decode_tool_result(&value)?),
_ => decode_text(&value)?,
};
let message_id = if runtime_event_id.is_some() || is_request_context {
wire_id
} else {
internal_id
};
Ok(CanonicalMessage {
message_id,
role,
origin,
content,
runtime_event_id,
})
}
pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
let wire: Value = serde_json::from_str(value)?;
let started_at_ms = wire
.pointer("/providerOptions/cursor/pendingToolCallStartedAtMs")
.and_then(Value::as_u64)
.ok_or_else(|| {
Error::Protocol("Cursor pending assistant is missing pendingToolCallStartedAtMs".into())
})?;
let internal_id = format!(
"cursor-pending:{}",
BlobId::digest(value.as_bytes()).to_base64()
);
let message = decode(value.as_bytes(), internal_id.clone())?;
let MessageContent::Assistant {
text,
thinking,
tool_round_id: _,
replay_state,
tool_calls,
} = message.content
else {
return Err(Error::Protocol(
"Cursor pending message is not an assistant message".into(),
));
};
if tool_calls.is_empty() {
return Err(Error::Protocol(
"Cursor resume contains a pending assistant without tool calls".into(),
));
}
let model_call_id = wire
.pointer("/providerOptions/cursor/modelProviderMessageId")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.unwrap_or(&internal_id)
.to_string();
let calls = tool_calls
.into_iter()
.enumerate()
.map(|(index, call)| {
Ok(ToolCall {
index,
call_id: call.call_id,
model_call_id: model_call_id.clone(),
name: call.name,
arguments_text: serde_json::to_string(&call.arguments)?,
arguments: call.arguments,
})
})
.collect::<Result<Vec<_>>>()?;
Ok(RecoveredToolRound {
assistant: ToolRoundAssistant {
text,
thinking,
model_call_id,
replay_state,
},
calls,
started_at_ms,
})
}
fn decode_text(value: &Value) -> Result<MessageContent> {
let content = value.get("content").unwrap_or(&Value::Null);
if let Some(text) = content.as_str() {
return Ok(MessageContent::Parts {
parts: vec![ContentPart::Text { text: text.into() }],
});
}
let parts = content
.as_array()
.ok_or_else(|| Error::Protocol("Cursor message content is not an array".into()))?
.iter()
.map(|part| match part.get("type").and_then(Value::as_str) {
Some("text") => Ok(ContentPart::Text {
text: required_string(part, "text")?.into(),
}),
Some("image") => {
let mime_type = required_string(part, "mimeType")?;
let encoded = required_string(part, "image")?;
Ok(ContentPart::Image {
mime_type: mime_type.into(),
data: STANDARD.decode(encoded).map_err(|error| {
Error::Protocol(format!("invalid Cursor image base64: {error}"))
})?,
})
}
Some(kind) => Err(Error::Protocol(format!(
"unsupported Cursor message content part: {kind}"
))),
None => Err(Error::Protocol(
"Cursor message content part is missing type".into(),
)),
})
.collect::<Result<Vec<_>>>()?;
Ok(MessageContent::Parts { parts })
}
fn decode_assistant(value: &Value, internal_id: &str) -> Result<MessageContent> {
let mut text = String::new();
let mut thinking = String::new();
let mut calls = Vec::new();
let mut replay_state = None;
for part in value
.get("content")
.and_then(Value::as_array)
.into_iter()
.flatten()
{
match part.get("type").and_then(Value::as_str) {
Some("text") => {
text.push_str(part.get("text").and_then(Value::as_str).unwrap_or_default())
}
Some("reasoning") => {
thinking.push_str(part.get("text").and_then(Value::as_str).unwrap_or_default());
if let Some(signature) = part.get("signature").and_then(Value::as_str) {
if replay_state.is_some() {
return Err(Error::Protocol(
"Cursor assistant has multiple reasoning signatures".into(),
));
}
replay_state = Some(decode_replay_state(signature)?);
}
}
Some("tool-call") => calls.push(ToolCallContent {
index: calls.len(),
call_id: required_string(part, "toolCallId")?.into(),
name: required_string(part, "toolName")?.into(),
arguments: part.get("args").cloned().unwrap_or(Value::Null),
}),
_ => {}
}
}
Ok(MessageContent::Assistant {
text,
thinking,
tool_round_id: (!calls.is_empty())
.then(|| ToolRoundId::new(format!("{internal_id}:tool-round"))),
replay_state,
tool_calls: calls,
})
}
fn decode_replay_state(signature: &str) -> Result<crate::model::ProviderReplayState> {
let Some(encoded) = signature.strip_prefix(REPLAY_ENVELOPE_PREFIX) else {
return Ok(crate::model::ProviderReplayState {
provider_kind: "cursor_opaque".into(),
value: Value::String(signature.into()),
});
};
let bytes = STANDARD.decode(encoded).map_err(|error| {
Error::Protocol(format!(
"invalid Cursor BYOK replay envelope base64: {error}"
))
})?;
serde_json::from_slice(&bytes)
.map_err(|error| Error::Protocol(format!("invalid Cursor BYOK replay envelope: {error}")))
}
fn decode_tool_result(value: &Value) -> Result<ToolResultContent> {
let part = value
.get("content")
.and_then(Value::as_array)
.and_then(|parts| parts.first())
.ok_or_else(|| Error::Protocol("Cursor tool message has no result part".into()))?;
Ok(ToolResultContent {
call_id: required_string(part, "toolCallId")?.into(),
name: required_string(part, "toolName")?.into(),
content: part
.get("result")
.and_then(Value::as_str)
.unwrap_or_default()
.into(),
is_error: part
.get("isError")
.and_then(Value::as_bool)
.unwrap_or(false),
image: None,
provider_parts: Vec::new(),
})
}
fn required_string<'a>(value: &'a Value, name: &str) -> Result<&'a str> {
value
.get(name)
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol(format!("Cursor message is missing {name}")))
}
@@ -0,0 +1,264 @@
//! Encodes canonical Messages into stable Cursor checkpoint message data.
use std::collections::HashSet;
use base64::{engine::general_purpose::STANDARD, Engine};
use serde_json::{json, Map, Value};
use crate::{
model::{
project_messages, CanonicalMessage, ContentPart, ProjectedContent, ProjectedMessage, Role,
ToolCall, ToolCallContent, ToolRoundAssistant,
},
Error, Result,
};
use super::REPLAY_ENVELOPE_PREFIX;
pub fn stable_messages(
instructions: &str,
messages: &[CanonicalMessage],
model: &str,
) -> Result<Vec<Vec<u8>>> {
let mut projected = project_messages(messages)?;
if !instructions.is_empty() {
projected.insert(
0,
ProjectedMessage {
message_id: "system".into(),
role: Role::System,
content: ProjectedContent::Parts(vec![ContentPart::Text {
text: instructions.into(),
}]),
},
);
}
projected
.iter()
.map(|message| serde_json::to_vec(&wire_message(message, model, None)?).map_err(Into::into))
.collect::<std::result::Result<_, _>>()
}
pub fn staged_tool_round(
assistant: &ToolRoundAssistant,
calls: &[ToolCall],
model: &str,
allowed_tools: &[String],
dynamic_tools: &HashSet<String>,
started_at_ms: u64,
) -> Result<String> {
let message = ProjectedMessage {
message_id: assistant.model_call_id.clone(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: assistant.text.clone(),
thinking: assistant.thinking.clone(),
replay_state: assistant.replay_state.clone(),
calls: calls
.iter()
.map(|call| ToolCallContent {
index: call.index,
call_id: call.call_id.clone(),
name: call.name.clone(),
arguments: call.arguments.clone(),
})
.collect(),
},
};
Ok(serde_json::to_string(&wire_message(
&message,
model,
Some(PendingContext {
allowed_tools,
dynamic_tools,
started_at_ms,
}),
)?)?)
}
pub fn staged_final(
message: &CanonicalMessage,
model: &str,
allowed_tools: &[String],
dynamic_tools: &HashSet<String>,
started_at_ms: u64,
) -> Result<String> {
let projected = project_messages(std::slice::from_ref(message))?;
let assistant = projected
.first()
.filter(|message| message.role == Role::Assistant)
.ok_or_else(|| {
Error::Protocol("final checkpoint stage is not an assistant message".into())
})?;
Ok(serde_json::to_string(&wire_message(
assistant,
model,
Some(PendingContext {
allowed_tools,
dynamic_tools,
started_at_ms,
}),
)?)?)
}
#[derive(Clone, Copy)]
pub(super) struct PendingContext<'a> {
allowed_tools: &'a [String],
dynamic_tools: &'a HashSet<String>,
started_at_ms: u64,
}
pub(super) fn wire_message(
message: &ProjectedMessage,
model: &str,
pending: Option<PendingContext<'_>>,
) -> Result<Value> {
let mut root = Map::new();
root.insert(
"role".into(),
Value::String(role_name(&message.role).into()),
);
root.insert("content".into(), wire_content(&message.content, model)?);
root.insert("id".into(), Value::String(wire_message_id(message)));
if let ProjectedContent::Assistant { calls, .. } = &message.content {
let mut cursor = Map::new();
if let Some(pending) = pending {
cursor.insert(
"pendingToolCallStartedAtMs".into(),
json!(pending.started_at_ms),
);
cursor.insert(
"pendingToolExecutionContracts".into(),
Value::Object(
calls
.iter()
.map(|call| {
(
call.call_id.clone(),
json!({
"toolCallId": call.call_id,
"outerToolName": call.name,
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
"isDynamic": pending.dynamic_tools.contains(&call.name),
"allowedToolNames": pending.allowed_tools,
}),
)
})
.collect(),
),
);
}
if !cursor.is_empty() {
root.insert("providerOptions".into(), json!({"cursor": cursor}));
}
}
Ok(Value::Object(root))
}
fn tool_identifier(name: &str, dynamic_tools: &HashSet<String>) -> String {
if dynamic_tools.contains(name) {
return name.into();
}
match name {
"CallMcpTool" | "SembleSearch" | "SembleFindRelated" => "MCP".into(),
"CreatePlan" => "CREATE_PLAN_V2".into(),
"UpdateCurrentStep" => "COMMUNICATE_UPDATE".into(),
_ => name
.chars()
.enumerate()
.fold(String::new(), |mut value, (index, character)| {
if index > 0 && character.is_ascii_uppercase() {
value.push('_');
}
value.push(character.to_ascii_uppercase());
value
}),
}
}
fn wire_message_id(message: &ProjectedMessage) -> String {
match &message.content {
ProjectedContent::Assistant { .. } => "1".into(),
ProjectedContent::ToolResult(result) => result.call_id.clone(),
ProjectedContent::Parts(_) => message.message_id.clone(),
}
}
fn wire_content(content: &ProjectedContent, model: &str) -> Result<Value> {
Ok(match content {
ProjectedContent::Parts(parts) => Value::Array(
parts
.iter()
.map(|part| match part {
ContentPart::Text { text } => json!({"type":"text", "text":text}),
ContentPart::Image { mime_type, data } => json!({
"type":"image",
"image": STANDARD.encode(data),
"mimeType": mime_type,
}),
})
.collect(),
),
ProjectedContent::Assistant {
text,
thinking,
replay_state,
calls,
} => {
let mut parts = Vec::new();
if !thinking.is_empty() || replay_state.is_some() {
let mut reasoning = json!({
"type": "reasoning",
"text": thinking,
"providerOptions": {"cursor": {"modelName": model}},
});
if let Some(replay_state) = replay_state {
reasoning["signature"] = Value::String(encode_replay_state(replay_state)?);
}
parts.push(reasoning);
}
if !text.is_empty() {
parts.push(json!({"type":"text", "text":text}));
}
parts.extend(calls.iter().map(|call| {
json!({
"type": "tool-call",
"toolCallId": call.call_id,
"toolName": call.name,
"args": call.arguments,
})
}));
Value::Array(parts)
}
ProjectedContent::ToolResult(result) => json!([{
"type": "tool-result",
"toolCallId": result.call_id,
"toolName": result.name,
"result": result.content,
"experimental_content": [{"type":"text", "text":result.content}],
"isError": result.is_error,
}]),
})
}
fn encode_replay_state(replay_state: &crate::model::ProviderReplayState) -> Result<String> {
if replay_state.provider_kind == "cursor_opaque" {
return replay_state
.value
.as_str()
.map(str::to_string)
.ok_or_else(|| Error::Protocol("Cursor opaque replay state is not a string".into()));
}
Ok(format!(
"{REPLAY_ENVELOPE_PREFIX}{}",
STANDARD.encode(serde_json::to_vec(replay_state)?)
))
}
fn role_name(role: &Role) -> &'static str {
match role {
Role::System => "system",
Role::User => "user",
Role::Assistant => "assistant",
Role::Tool => "tool",
}
}
@@ -0,0 +1,12 @@
//! Converts between canonical Messages and Cursor checkpoint message data.
mod decode;
mod encode;
pub use decode::{decode, decode_pending};
pub use encode::{stable_messages, staged_final, staged_tool_round};
const REPLAY_ENVELOPE_PREFIX: &str = "cursor-byok:v1:";
#[cfg(test)]
mod tests;
@@ -0,0 +1,261 @@
//! Verifies stable checkpoint Message encoding and recovery behavior.
use std::collections::HashSet;
use serde_json::{json, Value};
use crate::model::{
project_messages, CanonicalMessage, ContentPart, MessageContent, ProjectedContent,
ProjectedMessage, ProviderReplayState, Role, ToolCall, ToolResultContent, ToolRoundAssistant,
};
use super::{decode, decode_pending, encode::wire_message, staged_tool_round};
#[test]
fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
let replay_state = ProviderReplayState {
provider_kind: "anthropic".into(),
value: json!({"blocks":[{"type":"thinking","thinking":"why","signature":"sig"}]}),
};
let assistant = ToolRoundAssistant {
text: "before tools".into(),
thinking: "why".into(),
model_call_id: "model-call".into(),
replay_state: Some(replay_state.clone()),
};
let calls = vec![
ToolCall {
index: 0,
call_id: "a".into(),
model_call_id: "model-call".into(),
name: "Read".into(),
arguments_text: r#"{"path":"/a"}"#.into(),
arguments: json!({"path":"/a"}),
},
ToolCall {
index: 1,
call_id: "b".into(),
model_call_id: "model-call".into(),
name: "Grep".into(),
arguments_text: r#"{"pattern":"x"}"#.into(),
arguments: json!({"pattern":"x"}),
},
];
let pending = staged_tool_round(
&assistant,
&calls,
"claude",
&["Read".into(), "Grep".into()],
&HashSet::new(),
42,
)
.unwrap();
let wire: Value = serde_json::from_str(&pending).unwrap();
assert_eq!(wire["id"], "1");
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
"READ"
);
assert_eq!(wire["role"], "assistant");
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
.as_object()
.unwrap()
.len(),
2
);
assert_eq!(
wire["content"]
.as_array()
.unwrap()
.iter()
.filter(|part| part["type"] == "tool-call")
.count(),
2
);
let recovered = decode_pending(&pending).unwrap();
assert_eq!(recovered.assistant.replay_state, Some(replay_state));
assert_eq!(recovered.calls.len(), 2);
assert_eq!(recovered.calls[0].call_id, "a");
assert_eq!(recovered.calls[1].call_id, "b");
}
#[test]
fn cursor_wire_ids_are_projection_metadata_not_internal_message_ids() {
let assistant = ProjectedMessage {
message_id: "internal-assistant-id".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: "done".into(),
thinking: String::new(),
replay_state: None,
calls: Vec::new(),
},
};
let result = ProjectedMessage {
message_id: "internal-result-id".into(),
role: Role::Tool,
content: ProjectedContent::ToolResult(ToolResultContent {
call_id: "call-1".into(),
name: "Read".into(),
content: "ok".into(),
is_error: false,
image: None,
provider_parts: Vec::new(),
}),
};
assert_eq!(wire_message(&assistant, "model", None).unwrap()["id"], "1");
assert_eq!(
wire_message(&result, "model", None).unwrap()["id"],
"call-1"
);
}
#[test]
fn runtime_wire_identity_survives_checkpoint_hydration() {
let wire = json!({
"role": "user",
"id": "runtime:subagent-completed:child-id",
"content": "child completed",
});
let message = decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
"cursor-root:blob-id:19".into(),
)
.unwrap();
assert_eq!(message.message_id, "runtime:subagent-completed:child-id");
assert_eq!(
message.runtime_event_id.as_deref(),
Some("subagent-completed:child-id")
);
}
#[test]
fn request_context_identity_survives_checkpoint_hydration() {
let wire = json!({
"role": "user",
"id": "request-context:digest",
"content": "<rules>current rules</rules>",
});
let message = decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
"cursor-root:blob-id:20".into(),
)
.unwrap();
assert_eq!(message.message_id, "request-context:digest");
assert_eq!(message.origin, crate::model::Origin::Prompt);
}
#[test]
fn cursor_user_image_uses_image_field() {
let wire = json!({
"role": "user",
"id": "user-image",
"content": [
{"type":"text", "text":"look"},
{"type":"image", "image":"AQID", "mimeType":"image/png"},
],
});
let message = decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
"cursor-root:user-image".into(),
)
.unwrap();
assert!(matches!(
&message.content,
MessageContent::Parts { parts }
if parts[1] == ContentPart::Image {
mime_type: "image/png".into(),
data: vec![1, 2, 3],
}
));
let projected = project_messages(&[message]).unwrap();
let encoded = wire_message(&projected[0], "model", None).unwrap();
assert_eq!(encoded["content"][1]["image"], "AQID");
assert!(encoded["content"][1].get("data").is_none());
}
#[test]
fn repeated_cursor_wire_ids_do_not_merge_distinct_tool_rounds() {
fn assistant(call_id: &str, internal_id: &str) -> CanonicalMessage {
let wire = json!({
"role": "assistant",
"id": "1",
"content": [{
"type": "tool-call",
"toolCallId": call_id,
"toolName": "Read",
"args": {"path": format!("/{call_id}")},
}],
});
decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
internal_id.into(),
)
.unwrap()
}
fn result(call_id: &str, internal_id: &str) -> CanonicalMessage {
let wire = json!({
"role": "tool",
"id": call_id,
"content": [{
"type": "tool-result",
"toolCallId": call_id,
"toolName": "Read",
"result": "ok",
}],
});
decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
internal_id.into(),
)
.unwrap()
}
let messages = vec![
assistant("a", "cursor-root:a"),
result("a", "cursor-root:a-result"),
assistant("b", "cursor-root:b"),
result("b", "cursor-root:b-result"),
];
assert_ne!(messages[0].message_id, messages[2].message_id);
let projected = project_messages(&messages).unwrap();
assert_eq!(projected.len(), 4);
assert!(matches!(
&projected[0].content,
ProjectedContent::Assistant { calls, .. } if calls[0].call_id == "a"
));
assert!(matches!(
&projected[2].content,
ProjectedContent::Assistant { calls, .. } if calls[0].call_id == "b"
));
}
#[test]
fn opaque_cursor_reasoning_signature_round_trips_without_decoding() {
let signature = "opaque-url-safe_signature-value";
let wire = json!({
"role": "assistant",
"id": "1",
"content": [{"type":"reasoning", "text":"", "signature":signature}],
});
let message = decode(
serde_json::to_vec(&wire).unwrap().as_slice(),
"cursor-root:opaque".into(),
)
.unwrap();
let MessageContent::Assistant { replay_state, .. } = &message.content else {
panic!("expected assistant");
};
assert_eq!(
replay_state.as_ref().unwrap().provider_kind,
"cursor_opaque"
);
let projected = project_messages(&[message]).unwrap();
let encoded = wire_message(&projected[0], "model", None).unwrap();
assert_eq!(encoded["content"][0]["signature"], signature);
}
+14
View File
@@ -0,0 +1,14 @@
//! Builds, publishes, and restores Cursor Conversation checkpoints.
mod builder;
mod derived;
pub mod messages;
mod recovery;
mod roots;
mod steps;
mod summary;
mod turns;
pub(crate) mod worker;
pub use builder::CheckpointBuilder;
pub use steps::{PendingSteps, StepBuffer};
+49
View File
@@ -0,0 +1,49 @@
//! Restores Conversation Messages and pending Tool state from a checkpoint.
use crate::{
cursor::{checkpoint::messages, protocol::proto::agent::v1 as pb},
model::CanonicalMessage,
store::BlobId,
Error, Result,
};
use super::CheckpointBuilder;
impl CheckpointBuilder {
pub async fn import_prefetched(&self, blobs: &[pb::PreFetchedBlob]) -> Result<()> {
for blob in blobs {
let expected = BlobId::from_bytes(&blob.id)?;
let actual = self.store.put_blob(&blob.value, &[]).await?;
if expected != actual {
return Err(Error::Protocol(format!(
"prefetched Blob hash mismatch: {}",
expected.to_base64()
)));
}
}
Ok(())
}
pub async fn hydrate_messages(
&self,
state: Option<&pb::ConversationStateStructure>,
) -> Result<Vec<CanonicalMessage>> {
let mut messages = Vec::new();
let Some(state) = state else {
return Ok(messages);
};
for (ordinal, raw_id) in state.root_prompt_messages_json.iter().enumerate() {
let id = BlobId::from_bytes(raw_id)?;
let Some(data) = self.sync.get(&id).await? else {
return Err(Error::Protocol(format!(
"missing message Blob {}",
id.to_base64()
)));
};
messages.push(messages::decode(
&data,
format!("cursor-root:{}:{ordinal}", id.to_base64()),
)?);
}
Ok(messages)
}
}
+114
View File
@@ -0,0 +1,114 @@
//! Maintains stable append-only Cursor root messages.
use crate::{cursor::checkpoint::messages, model::CanonicalMessage, store::BlobId, Error, Result};
use super::CheckpointBuilder;
#[derive(Clone)]
pub(super) struct RootFrontier {
pub(super) ids: Vec<BlobId>,
pub(super) generated: Vec<Vec<u8>>,
pub(super) base_count: usize,
}
impl CheckpointBuilder {
pub(super) async fn project_roots(
&mut self,
messages: &[CanonicalMessage],
) -> Result<Vec<BlobId>> {
let wire_messages = messages::stable_messages(&self.instructions, messages, &self.model)?;
self.ensure_roots()?;
let replacement = self
.roots
.as_ref()
.and_then(|roots| changed_system_root(roots, &wire_messages));
if let Some(message) = replacement {
let id = self.sync.persist(&message, &[]).await?;
self.roots
.as_mut()
.ok_or_else(|| Error::Protocol("Cursor root frontier was not initialized".into()))?
.ids[0] = id;
}
let roots = self
.roots
.as_mut()
.ok_or_else(|| Error::Protocol("Cursor root frontier was not initialized".into()))?;
if wire_messages.len() < roots.ids.len() {
return Err(Error::Protocol(format!(
"Cursor stable history shrank from {} to {} roots",
roots.ids.len(),
wire_messages.len()
)));
}
for (index, expected) in roots.generated.iter().enumerate() {
let wire_index = roots.base_count + index;
if wire_messages.get(wire_index) != Some(expected) {
return Err(Error::Protocol(format!(
"Cursor stable root changed at index {wire_index}"
)));
}
}
for message in wire_messages.iter().skip(roots.ids.len()) {
roots.ids.push(self.sync.persist(message, &[]).await?);
roots.generated.push(message.clone());
}
Ok(roots.ids.clone())
}
fn ensure_roots(&mut self) -> Result<()> {
if self.roots.is_some() {
return Ok(());
}
let ids = self
.base
.root_prompt_messages_json
.iter()
.map(|id| BlobId::from_bytes(id))
.collect::<Result<Vec<_>>>()?;
self.roots = Some(RootFrontier {
base_count: ids.len(),
ids,
generated: Vec::new(),
});
Ok(())
}
pub(super) async fn replace_roots(
&mut self,
messages: &[CanonicalMessage],
) -> Result<Vec<BlobId>> {
let wire_messages = messages::stable_messages(&self.instructions, messages, &self.model)?;
self.ensure_roots()?;
let previous_system = self
.roots
.as_ref()
.and_then(|roots| roots.ids.first())
.cloned();
let mut ids = Vec::with_capacity(wire_messages.len());
for (index, message) in wire_messages.iter().enumerate() {
if index == 0
&& previous_system
.as_ref()
.is_some_and(|id| *id == BlobId::digest(message))
{
ids.push(previous_system.clone().expect("checked system root"));
} else {
ids.push(self.sync.persist(message, &[]).await?);
}
}
self.roots = Some(RootFrontier {
base_count: ids.len(),
ids: ids.clone(),
generated: Vec::new(),
});
Ok(ids)
}
}
fn changed_system_root(roots: &RootFrontier, messages: &[Vec<u8>]) -> Option<Vec<u8>> {
roots
.ids
.first()
.zip(messages.first())
.filter(|(current, message)| **current != BlobId::digest(message))
.map(|(_, message)| message.clone())
}
+116
View File
@@ -0,0 +1,116 @@
//! Buffers Conversation steps that have not yet been persisted.
use std::time::Duration;
use crate::cursor::{protocol::proto::agent::v1 as pb, tools::tool_call_result::ToolCompletion};
#[derive(Default)]
pub struct PendingSteps {
pub steps: Vec<pb::ConversationStep>,
pub read_paths: Vec<String>,
}
#[derive(Default)]
pub struct StepBuffer {
steps: Vec<pb::ConversationStep>,
read_paths: Vec<String>,
text: String,
thinking: String,
}
impl StepBuffer {
pub fn text_delta(&mut self, delta: &str) {
self.text.push_str(delta);
}
pub fn finish_text(&mut self) {
if self.text.is_empty() {
return;
}
self.steps.push(pb::ConversationStep {
message: Some(pb::conversation_step::Message::AssistantMessage(
pb::AssistantMessage {
text: std::mem::take(&mut self.text),
},
)),
});
}
pub fn thinking_delta(&mut self, delta: &str) {
self.thinking.push_str(delta);
}
pub fn finish_thinking(&mut self, duration: Duration) {
if self.thinking.is_empty() {
return;
}
self.steps.push(pb::ConversationStep {
message: Some(pb::conversation_step::Message::ThinkingMessage(
pb::ThinkingMessage {
text: std::mem::take(&mut self.thinking),
duration_ms: duration.as_millis().min(u32::MAX as u128) as u32,
},
)),
});
}
pub fn tool_completed(&mut self, completion: &ToolCompletion) {
if let Some(pb::tool_call::Tool::ReadToolCall(read)) = &completion.tool_call().tool {
if matches!(
read.result
.as_ref()
.and_then(|result| result.result.as_ref()),
Some(pb::read_tool_result::Result::Success(_))
) {
if let Some(path) = read.args.as_ref().map(|args| &args.path) {
if !path.is_empty() && !self.read_paths.contains(path) {
self.read_paths.push(path.clone());
}
}
}
}
self.steps.push(pb::ConversationStep {
message: Some(pb::conversation_step::Message::ToolCall(
completion.tool_call().clone(),
)),
});
}
pub fn discard_model_output(&mut self) {
self.text.clear();
self.thinking.clear();
self.steps.retain(|step| {
!matches!(
step.message,
Some(
pb::conversation_step::Message::AssistantMessage(_)
| pb::conversation_step::Message::ThinkingMessage(_)
)
)
});
}
pub fn take(&mut self) -> PendingSteps {
PendingSteps {
steps: std::mem::take(&mut self.steps),
read_paths: std::mem::take(&mut self.read_paths),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
let mut buffer = StepBuffer::default();
buffer.text_delta("partial answer");
buffer.finish_text();
buffer.thinking_delta("partial reasoning");
buffer.finish_thinking(Duration::from_millis(25));
buffer.discard_model_output();
assert!(buffer.take().steps.is_empty());
}
}
+101
View File
@@ -0,0 +1,101 @@
//! Builds compacted checkpoint summary state.
use prost::Message;
use crate::{
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
model::CanonicalMessage,
store::{BlobEdge, BlobId},
Error, Result,
};
use super::CheckpointBuilder;
impl CheckpointBuilder {
pub async fn compacted(
&mut self,
messages: &[CanonicalMessage],
mode: i32,
summary: &str,
presentation: &PendingSteps,
) -> Result<pb::ConversationStateStructure> {
let summarized = self
.base
.root_prompt_messages_json
.iter()
.skip(1)
.map(|id| BlobId::from_bytes(id))
.collect::<Result<Vec<_>>>()?;
let root_ids = self.replace_roots(messages).await?;
let summary_message = root_ids
.last()
.filter(|_| root_ids.len() >= 2)
.ok_or_else(|| Error::Protocol("compaction produced no summary root".into()))?
.clone();
let summary_id = self
.sync
.persist(
&pb::ConversationSummary {
summary: summary.into(),
}
.encode_to_vec(),
&[],
)
.await?;
let archive = pb::ConversationSummaryArchive {
summarized_messages: summarized.iter().map(|id| id.as_bytes().to_vec()).collect(),
summary: summary.into(),
window_tail: 0,
summary_message: summary_message.as_bytes().to_vec(),
};
let mut edges = summarized
.iter()
.enumerate()
.map(|(index, child)| BlobEdge {
child: child.clone(),
field_name: format!("summarized_messages[{index}]"),
})
.collect::<Vec<_>>();
edges.push(BlobEdge {
child: summary_message,
field_name: "summary_message".into(),
});
let archive_id = self.sync.persist(&archive.encode_to_vec(), &edges).await?;
let turn_ids = self.project_turns(mode, presentation).await?;
for path in &presentation.read_paths {
if !self.base.read_paths.contains(path) {
self.base.read_paths.push(path.clone());
}
}
self.base.root_prompt_messages_json =
root_ids.iter().map(|id| id.as_bytes().to_vec()).collect();
self.base.turns = turn_ids.iter().map(|id| id.as_bytes().to_vec()).collect();
self.base.pending_tool_calls.clear();
self.base.mode = Some(mode);
self.base.summary = Some(summary_id.as_bytes().to_vec());
self.base.summary_archive = Some(archive_id.as_bytes().to_vec());
if !self
.base
.summary_archives
.contains(&archive_id.as_bytes().to_vec())
{
self.base
.summary_archives
.push(archive_id.as_bytes().to_vec());
}
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
if let Some(details) = self.base.token_details.as_mut() {
details.breakdown = Some(crate::cursor::services::usage::breakdown(
details.used_tokens,
details.max_tokens,
details.breakdown.as_ref(),
&self.instructions,
&self.tool_definitions,
&self.dynamic_tools,
messages,
)?);
}
Ok(self.base.clone())
}
}
+127
View File
@@ -0,0 +1,127 @@
//! Projects buffered steps into Cursor Conversation turns.
use prost::Message;
use crate::{
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
store::{BlobEdge, BlobId},
Error, Result,
};
use super::CheckpointBuilder;
#[derive(Clone)]
pub(super) struct TurnFrontier {
pub(super) preceding: Vec<BlobId>,
pub(super) current_id: Option<BlobId>,
pub(super) current: pb::AgentConversationTurnStructure,
}
impl CheckpointBuilder {
pub(super) async fn project_turns(
&mut self,
mode: i32,
presentation: &PendingSteps,
) -> Result<Vec<BlobId>> {
self.ensure_turn(mode).await?;
let Some(turn) = self.turn.as_mut() else {
return self
.base
.turns
.iter()
.map(|id| BlobId::from_bytes(id))
.collect();
};
let changed = !presentation.steps.is_empty();
for step in &presentation.steps {
let mut encoded = Vec::new();
step.encode(&mut encoded)?;
let id = self.sync.persist(&encoded, &[]).await?;
turn.current.steps.push(id.as_bytes().to_vec());
}
if changed || turn.current_id.is_none() {
let wrapper = pb::ConversationTurnStructure {
turn: Some(
pb::conversation_turn_structure::Turn::AgentConversationTurn(
turn.current.clone(),
),
),
};
let mut encoded = Vec::new();
wrapper.encode(&mut encoded)?;
let mut edges = Vec::with_capacity(turn.current.steps.len() + 1);
edges.push(BlobEdge {
child: BlobId::from_bytes(&turn.current.user_message)?,
field_name: "agent_conversation_turn.user_message".into(),
});
for (index, raw_id) in turn.current.steps.iter().enumerate() {
edges.push(BlobEdge {
child: BlobId::from_bytes(raw_id)?,
field_name: format!("agent_conversation_turn.steps[{index}]"),
});
}
turn.current_id = Some(self.sync.persist(&encoded, &edges).await?);
}
let mut ids = turn.preceding.clone();
ids.push(
turn.current_id
.clone()
.ok_or_else(|| Error::Protocol("Cursor current Turn has no BlobID".into()))?,
);
Ok(ids)
}
async fn ensure_turn(&mut self, mode: i32) -> Result<()> {
if self.turns_initialized {
return Ok(());
}
self.turns_initialized = true;
let base_ids = self
.base
.turns
.iter()
.map(|id| BlobId::from_bytes(id))
.collect::<Result<Vec<_>>>()?;
if let Some(mut user) = self.turn_user.clone() {
user.mode = mode;
let mut encoded = Vec::new();
user.encode(&mut encoded)?;
let user_id = self.sync.persist(&encoded, &[]).await?;
self.turn = Some(TurnFrontier {
preceding: base_ids,
current_id: None,
current: pb::AgentConversationTurnStructure {
user_message: user_id.as_bytes().to_vec(),
steps: Vec::new(),
request_id: Some(self.sync.request_id().into()),
encrypted_model: None,
dynamic_tool_count: None,
send_message_step_indices: Vec::new(),
},
});
return Ok(());
}
let Some((current_id, preceding)) = base_ids.split_last() else {
return Ok(());
};
let data = self.sync.get(current_id).await?.ok_or_else(|| {
Error::Protocol(format!(
"missing current Turn Blob {}",
current_id.to_base64()
))
})?;
let wrapper = pb::ConversationTurnStructure::decode(data.as_slice())?;
let Some(pb::conversation_turn_structure::Turn::AgentConversationTurn(current)) =
wrapper.turn
else {
return Err(Error::Protocol(
"current Cursor Turn is not an agent conversation turn".into(),
));
};
self.turn = Some(TurnFrontier {
preceding: preceding.to_vec(),
current_id: Some(current_id.clone()),
current,
});
Ok(())
}
}
+206
View File
@@ -0,0 +1,206 @@
//! Serializes checkpoint jobs and completes commit barriers.
use tokio::sync::{mpsc, oneshot};
use crate::{
cursor::{
checkpoint::PendingSteps, protocol::proto::agent::v1 as pb, transport::TransportHandle,
},
model::{CheckpointId, ToolRoundId},
store::Store,
Error, Result,
};
use super::CheckpointBuilder;
pub(crate) struct CheckpointJob {
pub kind: CheckpointKind,
pub presentation: PendingSteps,
pub context_tokens: Option<u64>,
pub ready: Option<oneshot::Sender<std::result::Result<(), String>>>,
}
pub(crate) enum CheckpointKind {
Settled(CheckpointId),
ToolStarted {
round_id: ToolRoundId,
stable_checkpoint_id: CheckpointId,
},
ToolSettled(CheckpointId),
Final {
checkpoint_id: CheckpointId,
result: oneshot::Sender<Result<FinalCheckpoints>>,
},
Compaction {
checkpoint_id: CheckpointId,
summary: String,
result: oneshot::Sender<Result<pb::ConversationStateStructure>>,
},
}
pub(crate) struct FinalCheckpoints {
pub staged: pb::ConversationStateStructure,
pub settled: pb::ConversationStateStructure,
}
pub(crate) struct CheckpointWorker {
pub jobs: mpsc::Sender<CheckpointJob>,
pub failures: mpsc::Receiver<Error>,
task: tokio::task::JoinHandle<()>,
}
impl CheckpointWorker {
pub fn spawn(
store: Store,
mut builder: CheckpointBuilder,
handle: TransportHandle,
mode: i32,
) -> Self {
let (jobs, mut receiver) = mpsc::channel::<CheckpointJob>(32);
let (failures, failure_receiver) = mpsc::channel(1);
let task = tokio::spawn(async move {
while let Some(job) = receiver.recv().await {
builder.record_context_tokens(job.context_tokens);
let presentation = job.presentation;
let ready = job.ready;
let result = match job.kind {
CheckpointKind::Settled(checkpoint_id)
| CheckpointKind::ToolSettled(checkpoint_id) => {
publish_settled(
&store,
&mut builder,
&handle,
mode,
checkpoint_id,
&presentation,
)
.await
}
CheckpointKind::ToolStarted {
round_id,
stable_checkpoint_id,
} => {
publish_started(
&store,
&mut builder,
&handle,
mode,
round_id,
stable_checkpoint_id,
&presentation,
)
.await
}
CheckpointKind::Final {
checkpoint_id,
result,
} => {
let checkpoints =
build_final(&store, &mut builder, mode, checkpoint_id, &presentation)
.await;
let _ = result.send(checkpoints);
Ok(())
}
CheckpointKind::Compaction {
checkpoint_id,
summary,
result,
} => {
let messages = store.load_checkpoint_messages(checkpoint_id).await;
let checkpoint = match messages {
Ok(messages) => {
builder
.compacted(&messages, mode, &summary, &presentation)
.await
}
Err(error) => Err(error),
};
let _ = result.send(checkpoint);
Ok(())
}
};
if let Err(error) = result {
if let Some(ready) = ready {
let _ = ready.send(Err(error.to_string()));
}
tracing::error!(%error, "failed to build or publish Cursor checkpoint");
let _ = failures.send(error).await;
break;
}
if let Some(ready) = ready {
let _ = ready.send(Ok(()));
}
}
});
Self {
jobs,
failures: failure_receiver,
task,
}
}
pub fn abort(&self) {
self.task.abort();
}
}
async fn publish_settled(
store: &Store,
builder: &mut CheckpointBuilder,
handle: &TransportHandle,
mode: i32,
checkpoint_id: CheckpointId,
presentation: &PendingSteps,
) -> Result<()> {
let messages = store.load_checkpoint_messages(checkpoint_id).await?;
let checkpoint = builder.settled(&messages, mode, presentation).await?;
builder.publish(handle, &checkpoint).await
}
async fn publish_started(
store: &Store,
builder: &mut CheckpointBuilder,
handle: &TransportHandle,
mode: i32,
round_id: ToolRoundId,
stable_checkpoint_id: CheckpointId,
presentation: &PendingSteps,
) -> Result<()> {
let round = store
.tool_round(&round_id)
.await?
.ok_or_else(|| Error::Store(format!("checkpoint tool round not found: {round_id}")))?;
let messages = store.load_checkpoint_messages(stable_checkpoint_id).await?;
let checkpoint = builder
.staged_tool_round(
&messages,
mode,
&round.assistant,
&round.calls,
round.created_at_ms,
presentation,
)
.await?;
builder.publish(handle, &checkpoint).await
}
async fn build_final(
store: &Store,
builder: &mut CheckpointBuilder,
mode: i32,
checkpoint_id: CheckpointId,
presentation: &PendingSteps,
) -> Result<FinalCheckpoints> {
let messages = store.load_checkpoint_messages(checkpoint_id).await?;
let (assistant, stable) = messages
.split_last()
.ok_or_else(|| Error::Store("final checkpoint contains no assistant".into()))?;
let started_at_ms = crate::cursor::tools::runtime::now_ms();
let staged = builder
.staged_final(stable, mode, assistant, started_at_ms, presentation)
.await?;
let settled = builder
.settled(&messages, mode, &PendingSteps::default())
.await?;
Ok(FinalCheckpoints { staged, settled })
}
+54
View File
@@ -0,0 +1,54 @@
//! Classifies Cursor actions and selects their message delivery behavior.
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{CanonicalMessage, RunId},
};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessageDelivery {
Ignore,
InsertMessages,
BreakMessages,
}
#[derive(Clone, Debug)]
pub struct CompiledMessages {
pub event_id: String,
pub target_run_id: Option<RunId>,
pub messages: Vec<CanonicalMessage>,
pub delivery: MessageDelivery,
}
impl CompiledMessages {
pub fn ignored(event_id: impl Into<String>) -> Self {
Self {
event_id: event_id.into(),
target_run_id: None,
messages: Vec::new(),
delivery: MessageDelivery::Ignore,
}
}
}
pub fn delivery(action: &pb::conversation_action::Action) -> MessageDelivery {
use pb::conversation_action::Action;
match action {
Action::BackgroundTaskCompletionAction(_)
| Action::BackgroundShellAction(_)
| Action::BackgroundSubagentAction(_)
| Action::AsyncAskQuestionCompletionAction(_)
| Action::SubscriptionNotificationAction(_)
| Action::GoalContinuationAction(_) => MessageDelivery::InsertMessages,
Action::UserMessageAction(_) | Action::InjectContextAction(_) => {
MessageDelivery::BreakMessages
}
Action::CancelAction(_)
| Action::CancelSubagentAction(_)
| Action::ResumeAction(_)
| Action::SummarizeAction(_)
| Action::ShellCommandAction(_)
| Action::StartPlanAction(_)
| Action::ExecutePlanAction(_) => MessageDelivery::Ignore,
}
}
+392
View File
@@ -0,0 +1,392 @@
//! Compiles runtime information that interrupts the current cycle before appending.
use std::collections::BTreeMap;
use chrono::{Offset, Utc};
use chrono_tz::Tz;
use crate::{
cursor::{
prompting::{Mode, PromptCompiler},
protocol::proto::agent::v1 as pb,
services::blob_sync::BlobSynchronizer,
},
model::{CanonicalMessage, MessageContent, Origin, Role},
store::BlobId,
Error, Result,
};
use super::{context, images};
pub(crate) enum RuntimeAction {
Inject(pb::InjectContextAction),
UserMessage(pb::UserMessageAction),
}
pub(crate) async fn compile_user_message_action(
action: &pb::UserMessageAction,
current_mode: i32,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
let user = action
.user_message
.as_ref()
.ok_or_else(|| Error::Protocol("Cursor user message action has no UserMessage".into()))?;
if user.message_id.is_empty() {
return Err(Error::Protocol(
"Cursor user message action has no message_id".into(),
));
}
let mode = if user.mode == pb::AgentMode::Unspecified as i32 {
current_mode
} else {
user.mode
};
let mut action_context = action
.prepend_user_messages
.iter()
.map(|message| message.text.trim())
.filter(|text| !text.is_empty())
.map(str::to_string)
.collect::<Vec<_>>();
action_context.extend(
user.subagent_system_reminder
.iter()
.filter(|text| !text.is_empty())
.cloned(),
);
let empty_context = pb::RequestContext::default();
compile(
format!("user-message:{}", user.message_id),
super::run::mode_from_proto(mode)?,
user,
action.request_context.as_ref().unwrap_or(&empty_context),
&action_context.join("\\n\\n"),
compiler,
blobs,
)
.await
}
pub(crate) async fn compile_injection(
injection: &pb::InjectContextAction,
mode: i32,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
if injection.injection_id.is_empty() {
return Err(Error::Protocol(
"InjectContextAction has no injection_id".into(),
));
}
let event_id = format!("inject-context:{}", injection.injection_id);
match injection.payload.as_ref() {
Some(pb::inject_context_action::Payload::UserContext(context)) => {
let user = context.user_message.as_ref().ok_or_else(|| {
Error::Protocol("InjectContextAction UserContext has no UserMessage".into())
})?;
if user.message_id.is_empty() {
return Err(Error::Protocol(
"InjectContextAction UserMessage has no message_id".into(),
));
}
let empty_context = pb::RequestContext::default();
compile(
event_id,
super::run::mode_from_proto(mode)?,
user,
context.request_context.as_ref().unwrap_or(&empty_context),
"",
compiler,
blobs,
)
.await
}
Some(pb::inject_context_action::Payload::SystemContext(context)) => {
Ok(CanonicalMessage {
message_id: format!("runtime:{event_id}"),
role: Role::User,
origin: Origin::Runtime,
content: MessageContent::Parts {
parts: vec![crate::model::ContentPart::Text {
text: format!(
"<system_context_injection>\n<producer>{}</producer>\n{}\n</system_context_injection>",
context.producer, context.content
),
}],
},
runtime_event_id: Some(event_id),
})
}
None => Err(Error::Protocol(
"InjectContextAction has no payload".into(),
)),
}
}
pub async fn compile(
event_id: String,
mode: Mode,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
let timestamp = Time::now(
request_context
.env
.as_ref()
.map(|env| env.time_zone.as_str()),
)?
.timestamp;
compile_with_timestamp(
event_id,
mode,
user,
request_context,
action_context,
timestamp,
compiler,
blobs,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn user_event_id(
input_id: &str,
mode: Mode,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
projected_request_context: Option<&MessageContent>,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<String> {
let runtime = compile_with_timestamp(
"identity".into(),
mode,
user,
request_context,
action_context,
String::new(),
compiler,
blobs,
)
.await?;
let semantic = serde_json::to_vec(&(projected_request_context, runtime.content))?;
Ok(format!(
"{input_id}:{}",
BlobId::digest(&semantic).to_base64()
))
}
#[allow(clippy::too_many_arguments)]
async fn compile_with_timestamp(
event_id: String,
mode: Mode,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
timestamp: String,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
let mut values = BTreeMap::from([
("OPEN_FILES", section(open_files(user))),
(
"SELECTED_CONTEXT",
section(
context::selected_context(user)
.filter(|value| !value.is_empty())
.map(|value| format!("<selected_context>\n{value}\n</selected_context>"))
.unwrap_or_default(),
),
),
("ACTION_CONTEXT", section(action_context.to_string())),
("TIMESTAMP", timestamp),
("USER_QUERY", user.text.clone()),
("DEBUG_SERVER_ENDPOINT", String::new()),
("DEBUG_LOG_PATH", String::new()),
("DEBUG_SESSION_ID", String::new()),
]);
if let Some(debug) = &request_context.debug_mode_config {
values.insert("DEBUG_SERVER_ENDPOINT", debug.server_endpoint.clone());
values.insert("DEBUG_LOG_PATH", debug.log_path.clone());
values.insert("DEBUG_SESSION_ID", debug.session_id.clone());
}
message(
event_id,
user,
compiler.runtime_message(mode, &values)?,
blobs,
)
.await
}
pub(super) fn compile_request_context(
event_id: &str,
request_context: &pb::RequestContext,
history: &[CanonicalMessage],
) -> Result<Option<CanonicalMessage>> {
let time = Time::now(
request_context
.env
.as_ref()
.map(|env| env.time_zone.as_str()),
)?;
let text = context::compile_context(request_context, &time.today);
if text.is_empty() {
return Ok(None);
}
let message = CanonicalMessage::text(
format!("request-context:{event_id}"),
Role::User,
Origin::Prompt,
text,
);
Ok(should_project_request_context(history, &message).then_some(message))
}
fn should_project_request_context(
history: &[CanonicalMessage],
current: &CanonicalMessage,
) -> bool {
history
.iter()
.rev()
.find(|message| message.message_id.starts_with("request-context:"))
.is_none_or(|previous| previous.content != current.content)
}
pub async fn compile_background(
event_id: String,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
blobs: &BlobSynchronizer,
) -> Result<(CanonicalMessage, String)> {
let timestamp = Time::now(
request_context
.env
.as_ref()
.map(|env| env.time_zone.as_str()),
)?
.timestamp;
let text = format!(
"<timestamp>{timestamp}</timestamp>\n{}\n<user_query>{}</user_query>",
action_context.trim(),
user.text
);
let message = message(event_id, user, text.clone(), blobs).await?;
Ok((message, text))
}
async fn message(
event_id: String,
user: &pb::UserMessage,
text: String,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
Ok(CanonicalMessage {
message_id: format!("runtime:{event_id}"),
role: Role::User,
origin: Origin::Runtime,
content: MessageContent::Parts {
parts: images::parts(user, text, blobs).await?,
},
runtime_event_id: Some(event_id),
})
}
fn section(value: String) -> String {
let value = value.trim();
if value.is_empty() {
String::new()
} else {
format!("{value}\n\n")
}
}
fn open_files(user: &pb::UserMessage) -> String {
let Some(ide) = user
.selected_context
.as_ref()
.and_then(|selected| selected.invocation_context.as_ref())
.and_then(|invocation| invocation.data.as_ref())
.and_then(|data| match data {
pb::invocation_context::Data::IdeState(ide) => Some(ide),
_ => None,
})
else {
return String::new();
};
if ide.visible_files.is_empty() && ide.recently_viewed_files.is_empty() {
return String::new();
}
let mut output = String::from("<open_and_recently_viewed_files>\n");
if !ide.recently_viewed_files.is_empty() {
output.push_str("Recently viewed files (recent at the top, oldest at the bottom):\n");
for file in &ide.recently_viewed_files {
output.push_str(&format!(
"- {} (total lines: {})\n",
file.path, file.total_lines
));
}
output.push('\n');
}
if !ide.visible_files.is_empty() {
output.push_str("Files that are currently open and visible in the user's IDE:\n");
for (index, file) in ide.visible_files.iter().enumerate() {
output.push_str(&format!("- {} (", file.path));
if index == 0 {
output.push_str("currently focused file");
if let Some(cursor) = &file.cursor_position {
output.push_str(&format!(", cursor is on line {}", cursor.line));
}
output.push_str(&format!(", total lines: {}", file.total_lines));
} else {
output.push_str(&format!("total lines: {}", file.total_lines));
}
output.push_str(")\n");
}
output.push('\n');
}
output.push_str(
"Note: these files may or may not be relevant to the current conversation. Use the read file tool if you need to get the contents of some of them.\n</open_and_recently_viewed_files>",
);
output
}
struct Time {
timestamp: String,
today: String,
}
impl Time {
fn now(time_zone: Option<&str>) -> Result<Self> {
let zone = match time_zone.filter(|value| !value.is_empty()) {
Some(value) => value
.parse::<Tz>()
.map_err(|_| Error::Protocol(format!("invalid Cursor time zone: {value}")))?,
None => chrono_tz::UTC,
};
let now = Utc::now().with_timezone(&zone);
let offset = now.offset().fix().local_minus_utc();
let sign = if offset < 0 { '-' } else { '+' };
let offset = offset.unsigned_abs();
let hours = offset / 3600;
let minutes = (offset % 3600) / 60;
let utc = if minutes == 0 {
format!("UTC{sign}{hours}")
} else {
format!("UTC{sign}{hours}:{minutes:02}")
};
Ok(Self {
timestamp: format!("{} ({utc})", now.format("%A, %b %-d, %Y, %-I:%M %p")),
today: now.format("%A %b %-d,\n%Y").to_string(),
})
}
}
+565
View File
@@ -0,0 +1,565 @@
//! Compiles rules, skills, MCP metadata, and environment context.
use std::{
collections::{BTreeMap, HashMap, HashSet},
path::Path,
};
use prost::Message;
use serde_json::Value;
use crate::{
cursor::{
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
tools::runtime::McpRoute,
},
model::ToolDefinition,
store::BlobId,
Error, Result,
};
pub async fn hydrate(
request: &pb::AgentRunRequest,
context_sync: &RequestContextSynchronizer,
) -> Result<pb::RequestContext> {
let mut context = request_context(request).cloned().unwrap_or_default();
let Some(parts) = request
.action
.as_ref()
.and_then(|action| action.request_context_parts.as_ref())
else {
if is_background_completion(request) {
return context_sync
.load(request.conversation_id.as_deref().unwrap_or_default())
.await;
}
return Ok(context);
};
if let Some(current) = context_sync
.refresh_if_missing(
parts,
request.conversation_id.as_deref().unwrap_or_default(),
)
.await?
{
context.rules = current.rules;
context.non_file_rules = current.non_file_rules;
context.cloud_rule = current.cloud_rule;
context.agent_skills = current.agent_skills;
context.skill_options = current.skill_options;
context.custom_subagents = current.custom_subagents;
context.tools = current.tools;
context.mcp_instructions = current.mcp_instructions;
context.mcp_file_system_options = current.mcp_file_system_options;
context.mcp_meta_tool_options = current.mcp_meta_tool_options;
return Ok(context);
}
if let Some(part) = decode_part::<pb::RequestContextRulesPart>(
"rules",
&parts.rules_blob_id,
parts.rules_byte_length,
context_sync,
)
.await?
{
context.rules = part.rules;
context.non_file_rules = part.non_file_rules;
context.cloud_rule = part.cloud_rule;
}
if let Some(part) = decode_part::<pb::RequestContextSkillsPart>(
"skills",
&parts.skills_blob_id,
parts.skills_byte_length,
context_sync,
)
.await?
{
context.agent_skills = part.agent_skills;
context.skill_options = part.skill_options;
}
if let Some(part) = decode_part::<pb::RequestContextSubagentsPart>(
"subagents",
&parts.subagents_blob_id,
parts.subagents_byte_length,
context_sync,
)
.await?
{
context.custom_subagents = part.custom_subagents;
}
if let Some(part) = decode_part::<pb::RequestContextMcpsPart>(
"MCP",
&parts.mcps_blob_id,
parts.mcps_byte_length,
context_sync,
)
.await?
{
context.tools = part.tools;
context.mcp_instructions = part.mcp_instructions;
context.mcp_file_system_options = part.mcp_file_system_options;
context.mcp_meta_tool_options = part.mcp_meta_tool_options;
}
Ok(context)
}
fn is_background_completion(request: &pb::AgentRunRequest) -> bool {
matches!(
request
.action
.as_ref()
.and_then(|action| action.action.as_ref()),
Some(pb::conversation_action::Action::BackgroundTaskCompletionAction(_))
)
}
async fn decode_part<T: Message + Default>(
name: &str,
raw_id: &[u8],
expected_length: u32,
context_sync: &RequestContextSynchronizer,
) -> Result<Option<T>> {
if raw_id.is_empty() {
if expected_length != 0 {
return Err(Error::Protocol(format!(
"{name} context has a byte length but no BlobID"
)));
}
return Ok(None);
}
let id = BlobId::from_bytes(raw_id)?;
let data = context_sync.get(&id).await?.ok_or_else(|| {
Error::Protocol(format!(
"{name} context Blob is missing: {}",
id.to_base64()
))
})?;
if data.len() != expected_length as usize {
return Err(Error::Protocol(format!(
"{name} context Blob length mismatch: expected {expected_length}, got {}",
data.len()
)));
}
T::decode(data.as_slice())
.map(Some)
.map_err(|error| Error::Protocol(format!("invalid {name} context Blob: {error}")))
}
pub fn request_context(request: &pb::AgentRunRequest) -> Option<&pb::RequestContext> {
let action = request.action.as_ref()?;
action
.request_context_parts
.as_ref()
.and_then(|parts| parts.dynamic_context.as_ref())
.or_else(|| match action.action.as_ref()? {
pb::conversation_action::Action::UserMessageAction(action) => {
action.request_context.as_ref()
}
pb::conversation_action::Action::ExecutePlanAction(action) => {
action.request_context.as_ref()
}
_ => None,
})
}
pub fn compile_context(context: &pb::RequestContext, today: &str) -> String {
let mut sections = Vec::new();
let mut transcripts = None;
if let Some(env) = &context.env {
let workspace = env
.workspace_paths
.first()
.map(String::as_str)
.unwrap_or("");
let repo = context.git_repos.iter().find(|repo| repo.path == workspace);
sections.push(format!(
"<user_info>\nOS Version: {}\n\nShell: {}\n\nWorkspace Path: {}\n\nIs directory a git repo: {}\n\nTerminals folder: {}\n\nToday's date: {}\n\nNote: Prefer using absolute paths over relative paths as tool call args when possible.\n</user_info>",
env.os_version,
env.shell,
workspace,
repo.map(|repo| format!("Yes, at {}", repo.path)).unwrap_or_else(|| "No".into()),
env.terminals_folder,
today,
));
if !env.agent_transcripts_folder.is_empty() {
transcripts = Some(format!(
"<agent_transcripts>\nAgent transcripts (past chats) live in {}. They have names like <uuid>.jsonl, cite parent chat transcripts to the user as [<title for chat <=6 words>\n](<uuid excluding .jsonl>). Don't discuss the folder structure.\n</agent_transcripts>",
env.agent_transcripts_folder
));
}
}
sections.extend(context.git_repos.iter().map(|repo| {
format!(
"<git_status>\nThis is the git status at the start of the conversation. Note that this status is a snapshot in time, and will not update during the conversation.\n\n\nGit repo: {}\n\n```\n{}\n```\n</git_status>",
repo.path, repo.status
)
}));
sections.extend(transcripts);
let skill_contents = context
.agent_skills
.iter()
.map(|skill| skill.content.as_str())
.filter(|content| !content.is_empty())
.collect::<HashSet<_>>();
let mut rules = context
.rules
.iter()
.chain(context.non_file_rules.iter())
.filter(|rule| {
!rule.content.trim().is_empty()
&& !is_skill_rule(rule)
&& !skill_contents.contains(rule.content.as_str())
})
.map(|rule| format!("<user_rule>\n{}\n</user_rule>", rule.content))
.collect::<Vec<_>>();
rules.extend(
context
.cloud_rule
.iter()
.map(|rule| format!("<user_rule>\n{rule}\n</user_rule>")),
);
if !rules.is_empty() {
sections.push(format!("<rules>\n{}\n</rules>", rules.join("\n")));
}
let skills = context
.agent_skills
.iter()
.filter(|skill| !skill.disable_model_invocation)
.map(|skill| {
format!(
"<agent_skill fullPath=\"{}\">{}</agent_skill>",
xml(&skill.full_path),
xml(&skill.description),
)
})
.collect::<Vec<_>>();
if !skills.is_empty() {
sections.push(format!(
"<agent_skills>\n<available_skills>\n{}\n</available_skills>\n</agent_skills>",
skills.join("\n")
));
}
let subagents = context
.custom_subagents
.iter()
.map(|agent| {
format!(
"<subagent name=\"{}\">{}</subagent>",
xml(&agent.name),
agent.description
)
})
.collect::<Vec<_>>();
if !subagents.is_empty() {
sections.push(format!(
"<subagents>\n{}\n</subagents>",
subagents.join("\n")
));
}
{
let servers = context
.mcp_meta_tool_options
.as_ref()
.into_iter()
.flat_map(|options| &options.mcp_descriptors)
.filter_map(compile_mcp_descriptor)
.collect::<Vec<_>>();
if !servers.is_empty() {
sections.push(format!(
"<mcp_meta_tools>\nThe following MCP tools are available. Call a listed tool directly with CallMcpTool without calling GetMcpTools first. If a call returns an error, use it to correct the arguments or authentication and retry when appropriate.\n<mcp_meta_tool_servers>\n{}\n</mcp_meta_tool_servers>\n</mcp_meta_tools>",
servers.join("\n")
));
}
}
sections.join("\n\n")
}
fn compile_mcp_descriptor(server: &pb::McpDescriptor) -> Option<String> {
if server.server_identifier.trim().is_empty() {
return None;
}
let tools = server
.tools
.iter()
.filter(|tool| !tool.tool_name.trim().is_empty())
.map(|tool| {
let mut lines = vec![format!("<mcp_tool name=\"{}\">", xml(&tool.tool_name))];
if let Some(path) = tool
.definition_path
.as_deref()
.filter(|value| !value.trim().is_empty())
{
lines.push(format!("<definition_path>{}</definition_path>", xml(path)));
}
if let Some(description) = tool
.description
.as_deref()
.filter(|value| !value.trim().is_empty())
{
lines.push(format!("<description>{}</description>", xml(description)));
}
if let Some(schema) = mcp_input_schema(tool) {
lines.push(format!("<input_schema>{}</input_schema>", xml(&schema)));
}
lines.push("</mcp_tool>".into());
lines.join("\n")
})
.collect::<Vec<_>>();
if tools.is_empty() {
return None;
}
Some(format!(
"<mcp_meta_tool_server name=\"{}\" identifier=\"{}\">\n<tools>\n{}\n</tools>\n</mcp_meta_tool_server>",
xml(if server.server_name.trim().is_empty() {
&server.server_identifier
} else {
&server.server_name
}),
xml(&server.server_identifier),
tools.join("\n"),
))
}
fn mcp_input_schema(tool: &pb::McpToolDescriptor) -> Option<String> {
tool.input_schema_json
.as_deref()
.filter(|value| !value.trim().is_empty())
.map(|value| {
serde_json::from_str::<Value>(value)
.map(|value| value.to_string())
.unwrap_or_else(|_| value.to_string())
})
.or_else(|| {
tool.input_schema
.as_ref()
.map(prost_value)
.map(|value| value.to_string())
})
}
pub fn meta_mcp_routes(context: &pb::RequestContext) -> HashMap<(String, String), McpRoute> {
context
.mcp_meta_tool_options
.as_ref()
.into_iter()
.flat_map(|options| &options.mcp_descriptors)
.filter(|server| !server.server_identifier.trim().is_empty())
.flat_map(|server| {
server.tools.iter().filter_map(move |tool| {
if tool.tool_name.trim().is_empty() {
return None;
}
let provider_identifier = if server.server_name.trim().is_empty() {
server.server_identifier.clone()
} else {
server.server_name.clone()
};
Some((
(server.server_identifier.clone(), tool.tool_name.clone()),
McpRoute {
name: format!("{}-{}", server.server_identifier, tool.tool_name),
provider_identifier,
tool_name: tool.tool_name.clone(),
description: tool.description.clone().unwrap_or_default(),
},
))
})
})
.collect()
}
fn is_skill_rule(rule: &pb::CursorRule) -> bool {
Path::new(&rule.full_path)
.file_name()
.and_then(|name| name.to_str())
.is_some_and(|name| name.eq_ignore_ascii_case("SKILL.md"))
}
pub fn selected_context(user: &pb::UserMessage) -> Option<String> {
let selected = user.selected_context.as_ref()?;
let mut sections = selected.extra_context.clone();
sections.extend(
selected
.files
.iter()
.map(|file| format!("<file path=\"{}\">\n{}\n</file>", file.path, file.content)),
);
sections.extend(
selected
.code_selections
.iter()
.map(|value| format!("<code path=\"{}\">\n{}\n</code>", value.path, value.content)),
);
sections.extend(selected.terminals.iter().map(|value| {
format!(
"<terminal title=\"{}\">\n{}\n</terminal>",
value.title.as_deref().unwrap_or_default(),
value.content
)
}));
sections.extend(selected.terminal_selections.iter().map(|value| {
format!(
"<terminal_selection title=\"{}\">\n{}\n</terminal_selection>",
value.title.as_deref().unwrap_or_default(),
value.content
)
}));
sections.extend(selected.cursor_rules.iter().filter_map(|value| {
value.rule.as_ref().map(|rule| {
format!(
"<rule path=\"{}\">\n{}\n</rule>",
rule.full_path, rule.content
)
})
}));
sections.extend(selected.cursor_commands.iter().map(|value| {
format!(
"<command name=\"{}\">\n{}\n</command>",
value.name, value.content
)
}));
sections.extend(selected.selected_skills.iter().map(|value| {
format!(
"<skill path=\"{}\">\n{}\n{}\n</skill>",
value.full_path, value.description, value.content
)
}));
sections.extend(selected.external_links.iter().map(|value| {
format!(
"External link: {}{}",
value.url,
value
.pdf_content
.as_deref()
.map(|content| format!("\n{content}"))
.unwrap_or_default()
)
}));
Some(sections.join("\n\n"))
}
pub fn dynamic_mcp(
request: &pb::AgentRunRequest,
context: &pb::RequestContext,
) -> Result<BTreeMap<String, (pb::McpToolDefinition, ToolDefinition)>> {
let direct = request
.mcp_tools
.iter()
.flat_map(|tools| tools.mcp_tools.iter());
let contextual = context.tools.iter();
let mut output = BTreeMap::new();
for wire in direct.chain(contextual) {
if wire.name.is_empty() {
return Err(Error::Protocol(
"MCP tool definition is missing name".into(),
));
}
let parameters = match wire.input_schema_json.as_deref() {
Some(json) if !json.trim().is_empty() => serde_json::from_str(json)?,
_ => prost_value(wire.input_schema.as_ref().ok_or_else(|| {
Error::Protocol(format!("MCP tool {} is missing input schema", wire.name))
})?),
};
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
let name = model_tool_name(&wire.name);
let definition = ToolDefinition {
name: name.clone(),
description: wire.description.clone(),
parameters,
};
if output
.insert(name.clone(), (wire.clone(), definition))
.is_some()
{
return Err(Error::Protocol(format!(
"duplicate MCP tool name after normalization: {name}"
)));
}
}
Ok(output)
}
fn normalize_mcp_parameters(tool_name: &str, mut parameters: Value) -> Result<Value> {
let schema = parameters
.as_object_mut()
.ok_or_else(|| invalid_mcp_parameters(tool_name))?;
match schema.get("type") {
Some(Value::String(schema_type)) if schema_type == "object" => return Ok(parameters),
Some(_) => return Err(invalid_mcp_parameters(tool_name)),
None => {}
}
let object_only_union = ["anyOf", "oneOf"].into_iter().any(|keyword| {
schema
.get(keyword)
.and_then(Value::as_array)
.is_some_and(|branches| {
!branches.is_empty()
&& branches.iter().all(|branch| {
branch
.as_object()
.and_then(|branch| branch.get("type"))
.and_then(Value::as_str)
== Some("object")
})
})
});
if !object_only_union {
return Err(invalid_mcp_parameters(tool_name));
}
// OpenAI-compatible function schemas (and the corresponding schema
// validators used by other providers) require the root schema to declare
// an object type. Cursor's app-control MCP sometimes sends an object-only
// `anyOf`/`oneOf` schema without that root annotation. Preserve the union
// while adding the annotation to the model-facing copy.
schema.insert("type".into(), Value::String("object".into()));
Ok(parameters)
}
fn invalid_mcp_parameters(tool_name: &str) -> Error {
Error::Protocol(format!(
"MCP tool {tool_name} input schema must describe an object"
))
}
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() {
None | Some(Kind::NullValue(_)) => Value::Null,
Some(Kind::NumberValue(value)) => serde_json::Number::from_f64(*value)
.map(Value::Number)
.unwrap_or(Value::Null),
Some(Kind::StringValue(value)) => Value::String(value.clone()),
Some(Kind::BoolValue(value)) => Value::Bool(*value),
Some(Kind::StructValue(value)) => Value::Object(
value
.fields
.iter()
.map(|(key, value)| (key.clone(), prost_value(value)))
.collect(),
),
Some(Kind::ListValue(value)) => {
Value::Array(value.values.iter().map(prost_value).collect())
}
}
}
fn xml(value: &str) -> String {
value
.replace('&', "&amp;")
.replace('"', "&quot;")
.replace('<', "&lt;")
.replace('>', "&gt;")
}
+75
View File
@@ -0,0 +1,75 @@
//! Resolves and persists images and blobs referenced by Cursor inputs.
use crate::{
cursor::{protocol::proto::agent::v1 as pb, services::blob_sync::BlobSynchronizer},
model::ContentPart,
store::BlobId,
Error, Result,
};
pub async fn parts(
message: &pb::UserMessage,
text: String,
blobs: &BlobSynchronizer,
) -> Result<Vec<ContentPart>> {
let mut parts = vec![ContentPart::Text { text }];
if let Some(context) = &message.selected_context {
for image in &context.selected_images {
parts.push(ContentPart::Image {
mime_type: image_mime_type(image)?,
data: image_data(image, blobs).await?,
});
}
}
Ok(parts)
}
fn image_mime_type(image: &pb::SelectedImage) -> Result<String> {
let mime_type = image.mime_type.trim();
if !mime_type.starts_with("image/") || mime_type.len() == "image/".len() {
return Err(Error::Protocol(format!(
"selected image has invalid MIME type: {}",
image.mime_type
)));
}
Ok(mime_type.into())
}
async fn image_data(image: &pb::SelectedImage, blobs: &BlobSynchronizer) -> Result<Vec<u8>> {
use pb::selected_image::DataOrBlobId;
let data = match image.data_or_blob_id.as_ref() {
Some(DataOrBlobId::Data(data)) => data.clone(),
Some(DataOrBlobId::BlobId(raw_id)) => {
let id = BlobId::from_bytes(raw_id)?;
blobs.get(&id).await?.ok_or_else(|| {
Error::Protocol(format!(
"selected image Blob is missing: {}",
id.to_base64()
))
})?
}
Some(DataOrBlobId::BlobIdWithData(value)) => {
let id = BlobId::from_bytes(&value.blob_id)?;
if value.data.is_empty() {
blobs.get(&id).await?.ok_or_else(|| {
Error::Protocol(format!(
"selected image Blob is missing: {}",
id.to_base64()
))
})?
} else {
blobs.cache_received(&id, &value.data).await?;
value.data.clone()
}
}
None => {
return Err(Error::Protocol(
"selected image is missing data_or_blob_id".into(),
))
}
};
if data.is_empty() {
return Err(Error::Protocol("selected image data is empty".into()));
}
Ok(data)
}
@@ -0,0 +1,210 @@
//! Compiles non-interrupting runtime information into append-only Messages.
use std::collections::BTreeMap;
use crate::{cursor::protocol::proto::agent::v1 as pb, Error, Result};
pub(super) const FOLLOW_UP: &str = concat!(
"Perform any necessary follow-up actions in response to the subagent completion above. ",
"If no follow-up work is needed, no further action is required. ",
"If you mention an agent or subagent in your response, link it with the `[Name](id)` ",
"Don't use generic label such as `[agent]`, `[worker]`, or `[subagent]`. ",
"For cloud subagents, when the agent has edited code, link to `[Review](bc-id#changes)`, ",
"or, if you know the exact added and deleted line counts, `[Review +A −D](bc-id#changes)`, ",
"replacing A and D with those counts. Never write A or D literally. ",
"Use `[Try Live](bc-id#desktop)` only when the agent used computer use. ",
"Don't repeat the same confirmation every time."
);
pub(super) const SHELL_FOLLOW_UP: &str = concat!(
"Briefly inform the user about the task result and perform any follow-up actions (if needed). ",
"If there's no follow-ups needed, don't explicitly say that."
);
#[derive(Debug)]
pub(super) struct Projection {
pub context: String,
pub turn_user: pb::UserMessage,
}
pub(super) fn project(
action: &pb::BackgroundTaskCompletionAction,
mode: i32,
) -> Result<Projection> {
if action.completions.is_empty() {
return Err(Error::Protocol(
"background task completion action contains no completion".into(),
));
}
let mut completions = BTreeMap::new();
let mut has_shell = false;
let mut has_subagent = false;
for completion in &action.completions {
let kind = pb::BackgroundTaskKind::try_from(completion.kind).map_err(|_| {
Error::Protocol(format!("unknown background task kind: {}", completion.kind))
})?;
if kind == pb::BackgroundTaskKind::Unspecified {
return Err(Error::Protocol(format!(
"background task completion has invalid kind: {}",
kind.as_str_name()
)));
}
let reason =
pb::BackgroundTaskCompletionReason::try_from(completion.reason).map_err(|_| {
Error::Protocol(format!(
"unknown background task completion reason: {}",
completion.reason
))
})?;
if reason != pb::BackgroundTaskCompletionReason::TaskFinished {
// Progress and reparenting notifications are informational; the
// client batches them together with the real finish notification.
continue;
}
if completion.task_id.is_empty() || completion.title.is_empty() {
return Err(Error::Protocol(
"background task completion requires task_id and title".into(),
));
}
let agent_id = match kind {
pb::BackgroundTaskKind::Shell => {
has_shell = true;
None
}
pb::BackgroundTaskKind::Subagent => {
has_subagent = true;
Some(
completion
.subagent_id
.as_deref()
.filter(|id| !id.is_empty())
.ok_or_else(|| {
Error::Protocol(
"background subagent completion has no subagent_id".into(),
)
})?,
)
}
pb::BackgroundTaskKind::Unspecified => unreachable!(),
};
let tool_call_id = completion
.tool_call_id
.as_deref()
.filter(|id| !id.is_empty())
.ok_or_else(|| {
Error::Protocol("background task completion has no tool_call_id".into())
})?;
let task_identity = agent_id.unwrap_or(&completion.task_id);
let identity = format!("{}:{task_identity}:{tool_call_id}", kind.as_str_name());
let context = completion_context(completion, kind, agent_id)?;
if completions
.insert(identity.clone(), (completion, context))
.is_some()
{
return Err(Error::Protocol(format!(
"duplicate background task completion: {identity}"
)));
}
}
let (first, _) = completions.values().next().ok_or_else(|| {
Error::Protocol("background task notification contains no finished task".into())
})?;
let text = match (has_shell, has_subagent) {
(true, false) => SHELL_FOLLOW_UP.into(),
(false, true) => FOLLOW_UP.into(),
(true, true) => format!("{SHELL_FOLLOW_UP}\n\n{FOLLOW_UP}"),
(false, false) => unreachable!(),
};
Ok(Projection {
context: completions
.values()
.map(|(_, context)| context.as_str())
.collect::<Vec<_>>()
.join("\n\n"),
turn_user: pb::UserMessage {
text,
message_id: format!(
"background-completed:{}",
completions.keys().cloned().collect::<Vec<_>>().join(":")
),
mode,
is_simulated_msg: Some(true),
simulated_msg_reason: Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32),
simulated_message_metadata: Some(pb::user_message::SimulatedMessageMetadata {
title: Some(first.title.clone()),
task_id: Some(first.task_id.clone()),
..Default::default()
}),
..Default::default()
},
})
}
fn status(completion: &pb::BackgroundTaskCompletion) -> Result<pb::BackgroundTaskStatus> {
let status = pb::BackgroundTaskStatus::try_from(completion.status).map_err(|_| {
Error::Protocol(format!(
"unknown background task status: {}",
completion.status
))
})?;
if status == pb::BackgroundTaskStatus::Unspecified {
return Err(Error::Protocol(
"background task completion has unspecified status".into(),
));
}
Ok(status)
}
fn completion_context(
completion: &pb::BackgroundTaskCompletion,
kind: pb::BackgroundTaskKind,
agent_id: Option<&str>,
) -> Result<String> {
let status = status(completion)?;
let mut fields = vec![
format!(
"kind: {}",
match kind {
pb::BackgroundTaskKind::Shell => "shell",
pb::BackgroundTaskKind::Subagent => "subagent",
pb::BackgroundTaskKind::Unspecified => unreachable!(),
}
),
format!("status: {}", status_name(status)),
format!("task_id: {}", completion.task_id),
format!("title: {}", completion.title),
];
optional_field(
&mut fields,
"tool_call_id",
completion.tool_call_id.as_deref(),
);
optional_field(&mut fields, "agent_id", agent_id);
optional_field(&mut fields, "detail", completion.detail.as_deref());
optional_field(
&mut fields,
"output_path",
completion.output_path.as_deref(),
);
optional_field(&mut fields, "thread_id", completion.thread_id.as_deref());
Ok(format!(
"<system_notification>\nThe following task has finished. If you were already aware, ignore this notification and do not restate prior responses.\n\n<task>\n{}\n</task>\n</system_notification>",
fields.join("\n")
))
}
fn optional_field(fields: &mut Vec<String>, name: &str, value: Option<&str>) {
if let Some(value) = value.filter(|value| !value.is_empty()) {
fields.push(format!("{name}: {value}"));
}
}
fn status_name(status: pb::BackgroundTaskStatus) -> &'static str {
match status {
pb::BackgroundTaskStatus::Success => "success",
pb::BackgroundTaskStatus::Error => "error",
pb::BackgroundTaskStatus::Aborted => "aborted",
pb::BackgroundTaskStatus::Unspecified => unreachable!(),
}
}
+13
View File
@@ -0,0 +1,13 @@
//! Compiles Cursor requests and actions into provider-independent Run inputs.
mod action;
mod break_messages;
mod context;
mod images;
mod insert_messages;
mod model;
mod run;
pub use action::*;
pub(crate) use break_messages::{compile_injection, compile_user_message_action, RuntimeAction};
pub use run::*;
+141
View File
@@ -0,0 +1,141 @@
//! Resolves Cursor model selections to configured provider models.
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{
parse_token_count, ModelLatency, ModelSpec, ReasoningSpec, SubagentKind,
SubagentModelOverride,
},
Error, Result,
};
pub fn requested_model(request: &pb::AgentRunRequest) -> Result<ModelSpec> {
let details = request.model_details.as_ref();
let model = if let Some(requested) = request.requested_model.as_ref() {
from_requested(requested, details)?
} else if let Some(model_id) = details
.map(|model| model.model_id.as_str())
.filter(|model| !model.is_empty())
{
ModelSpec {
model_id: model_id.into(),
display_name: details
.map(|model| model.display_name.clone())
.filter(|name| !name.is_empty()),
reasoning: ReasoningSpec {
enabled: details.is_some_and(|model| model.thinking_details.is_some()),
effort: None,
},
latency: ModelLatency::Standard,
max_output_tokens: None,
context_window_tokens: None,
supports_image_generation: false,
extra_params: serde_json::json!({}),
}
} else {
return Err(Error::Protocol("Cursor Run does not select a model".into()));
};
Ok(model)
}
pub fn overrides(
request: &pb::AgentRunRequest,
) -> Result<Vec<(SubagentKind, SubagentModelOverride)>> {
request
.subagent_model_overrides
.iter()
.map(|value| {
use pb::subagent_model_override::Selection;
let kind = subagent_kind(&value.subagent_type);
let selection = match value.selection.as_ref() {
Some(Selection::Model(model)) => {
if model.model_id == "default" {
SubagentModelOverride::Inherit
} else {
SubagentModelOverride::Explicit(from_requested(model, None)?)
}
}
Some(Selection::Inherit(true)) => SubagentModelOverride::Inherit,
Some(Selection::Disabled(true)) => SubagentModelOverride::Disabled,
None | Some(Selection::Inherit(false) | Selection::Disabled(false)) => {
return Err(Error::Protocol(format!(
"Cursor subagent model override {} has no active selection",
value.subagent_type
)))
}
};
Ok((kind, selection))
})
.collect()
}
pub fn subagent_kind(value: &str) -> SubagentKind {
if value == "generalPurpose" {
SubagentKind::GeneralPurpose
} else {
SubagentKind::Named(value.into())
}
}
fn from_requested(
model: &pb::RequestedModel,
details: Option<&pb::ModelDetails>,
) -> Result<ModelSpec> {
let mut spec = ModelSpec {
model_id: model.model_id.clone(),
display_name: details
.map(|model| model.display_name.clone())
.filter(|name| !name.is_empty()),
reasoning: ReasoningSpec {
enabled: model.max_mode
|| details.is_some_and(|model| model.thinking_details.is_some()),
effort: None,
},
latency: ModelLatency::Standard,
max_output_tokens: None,
context_window_tokens: None,
supports_image_generation: false,
extra_params: serde_json::json!({}),
};
for parameter in &model.parameters {
match parameter.id.as_str() {
"effort" | "reasoning" => {
let effort = parameter.value.trim();
spec.reasoning.effort =
(effort != "none" && !effort.is_empty()).then(|| effort.to_string());
spec.reasoning.enabled |= spec.reasoning.effort.is_some();
}
"thinking" => spec.reasoning.enabled |= parse_bool(parameter)?,
"fast" => {
if parse_bool(parameter)? {
spec.latency = ModelLatency::Fast;
}
}
"context" => {
spec.context_window_tokens =
Some(parse_token_count(&parameter.value).ok_or_else(|| {
Error::Protocol(format!(
"invalid Cursor context token count: {}",
parameter.value
))
})?);
}
other => {
return Err(Error::Protocol(format!(
"unsupported Cursor model parameter: {other}"
)))
}
}
}
Ok(spec)
}
fn parse_bool(parameter: &pb::requested_model::ModelParameterValue) -> Result<bool> {
match parameter.value.as_str() {
"true" => Ok(true),
"false" => Ok(false),
_ => Err(Error::Protocol(format!(
"invalid Cursor boolean model parameter {}={}",
parameter.id, parameter.value
))),
}
}
+610
View File
@@ -0,0 +1,610 @@
//! Compiles an AgentRunRequest into a PreparedRun.
use std::collections::BTreeMap;
use uuid::Uuid;
use crate::{
cursor::prompting::{Mode, PromptCompiler},
cursor::{
checkpoint::messages,
checkpoint::CheckpointBuilder,
protocol::proto::agent::v1 as pb,
services::blob_sync::BlobSynchronizer,
services::context_sync::RequestContextSynchronizer,
tools::runtime::{ExecContext, SubagentModel},
},
model::{
CanonicalMessage, ContentPart, ConversationId, MessageContent, Origin, PreparedRun,
PromptSpec, Role, RunAction, RunId, RunKind,
},
store::{BlobId, Store},
Error, Result,
};
use super::{break_messages, context, insert_messages, model};
struct ActionProjection {
mode: i32,
turn_user: Option<pb::UserMessage>,
action_context: String,
event_id: Option<String>,
input_id: Option<String>,
starts_turn: bool,
compacting: bool,
background_completion: bool,
}
pub struct CursorRunContext {
pub request_id: String,
pub mode: i32,
pub turn_user: Option<pb::UserMessage>,
pub exec: ExecContext,
pub dynamic_tools: BTreeMap<String, pb::McpToolDefinition>,
pub checkpoint_prompt: PromptSpec,
pub compacting: bool,
pub background_completion: bool,
}
pub(crate) struct PrepareDependencies<'a> {
pub compiler: &'a PromptCompiler,
pub store: &'a Store,
pub checkpoint: &'a CheckpointBuilder,
pub blob_sync: &'a BlobSynchronizer,
pub context_sync: &'a RequestContextSynchronizer,
}
pub(crate) async fn prepare(
request_id: &str,
request: &pb::AgentRunRequest,
dependencies: PrepareDependencies<'_>,
) -> Result<(PreparedRun, CursorRunContext)> {
let PrepareDependencies {
compiler,
store,
checkpoint,
blob_sync,
context_sync,
} = dependencies;
checkpoint
.import_prefetched(&request.pre_fetched_blobs)
.await?;
let conversation_id = ConversationId::new(
request
.conversation_id
.clone()
.unwrap_or_else(|| request_id.into()),
);
let run_id = execution_run_id(request_id);
let mut base_messages = if request.conversation_state.is_some() {
Some(
checkpoint
.hydrate_messages(request.conversation_state.as_ref())
.await?,
)
} else {
None
};
if let Some(trace) = blob_sync.trace() {
let hydrated_messages = base_messages.as_deref().unwrap_or_default();
let hydrated_images = hydrated_messages
.iter()
.map(|message| match &message.content {
MessageContent::Parts { parts } => parts
.iter()
.filter(|part| matches!(part, ContentPart::Image { .. }))
.count(),
_ => 0,
})
.sum::<usize>();
let history = request
.action
.as_ref()
.and_then(|action| action.action.as_ref())
.and_then(|action| match action {
pb::conversation_action::Action::UserMessageAction(action) => {
action.conversation_history.as_ref()
}
_ => None,
});
let summary = serde_json::json!({
"checkpoint_root_count": request.conversation_state.as_ref().map_or(0, |state| state.root_prompt_messages_json.len()),
"checkpoint_turn_count": request.conversation_state.as_ref().map_or(0, |state| state.turns.len()),
"conversation_history_message_count": history.map_or(0, |history| history.messages.len()),
"hydrated_message_count": hydrated_messages.len(),
"hydrated_image_count": hydrated_images,
"selected_source": "root_prompt_messages_json",
});
let encoded = serde_json::to_vec(&summary)?;
trace
.artifact("history_projection", "byok_server", &encoded, summary)
.await;
}
let request_context = context::hydrate(request, context_sync).await?;
let ActionProjection {
mode: mode_number,
mut turn_user,
action_context,
mut event_id,
input_id,
starts_turn,
compacting,
background_completion,
} = action(request)?;
let checkpoint_mode = if request.subagent_type_name.is_some() {
Mode::Subagent
} else {
mode_from_proto(mode_number)?
};
let mut model = model::requested_model(request)?;
if let Some(configured_model) = store.model(&model.model_id).await? {
configured_model.configure(&mut model);
}
let dynamic = context::dynamic_mcp(request, &request_context)?;
let subagent_model_overrides = model::overrides(request)?;
let subagents_disabled = subagent_model_overrides
.first()
.is_some_and(|(_, selection)| {
matches!(selection, crate::model::SubagentModelOverride::Disabled)
});
let mut checkpoint_prompt = compiler.prompt_spec(
checkpoint_mode,
&model,
&dynamic
.values()
.map(|(_, definition)| definition.clone())
.collect::<Vec<_>>(),
request.suppress_subagent_progress_update_tool == Some(true),
)?;
if subagents_disabled {
checkpoint_prompt.tools.retain(|tool| tool.name != "Task");
}
let prompt = if compacting {
compiler.prompt_spec(Mode::Compaction, &model, &[], false)?
} else {
checkpoint_prompt.clone()
};
let proposed_base_checkpoint_id = match base_messages.as_mut() {
Some(messages) if !messages.is_empty() => {
validate_prompt_root(messages)?;
messages.retain(|message| {
!(message.role == Role::System && message.origin == Origin::Prompt)
});
store.import_checkpoint(&conversation_id, messages).await?
}
Some(_) | None => store.ensure_conversation(&conversation_id).await?,
};
let base_checkpoint_id = match input_id.as_deref() {
Some(input_id) => {
store
.anchor_input(&conversation_id, input_id, proposed_base_checkpoint_id)
.await?
}
None => proposed_base_checkpoint_id,
};
let mut projected_user_context = if input_id.is_some() && !compacting && !background_completion
{
break_messages::compile_request_context(
"identity",
&request_context,
base_messages.as_deref().unwrap_or_default(),
)?
} else {
None
};
if event_id.is_none() {
if let (Some(input_id), Some(user)) = (input_id.as_deref(), turn_user.as_ref()) {
event_id = Some(
break_messages::user_event_id(
input_id,
checkpoint_mode,
user,
&request_context,
&action_context,
projected_user_context
.as_ref()
.map(|message| &message.content),
compiler,
blob_sync,
)
.await?,
);
}
}
let existing_runtime = match event_id.as_deref() {
Some(event_id) => {
store
.message(&conversation_id, &format!("runtime:{event_id}"))
.await?
}
_ => None,
};
let request_context_message = match event_id.as_deref() {
Some(event_id) if !compacting && !background_completion => {
let message_id = format!("request-context:{event_id}");
match store.message(&conversation_id, &message_id).await? {
Some(message) => Some(message),
None if input_id.is_some() => projected_user_context.take().map(|mut message| {
message.message_id = message_id;
message
}),
None => break_messages::compile_request_context(
event_id,
&request_context,
base_messages.as_deref().unwrap_or_default(),
)?,
}
}
_ => None,
};
let mut initial_messages = if compacting {
Vec::new()
} else {
match (turn_user.clone(), event_id) {
(Some(mut user), Some(event_id)) if background_completion => {
let (message, text) = match existing_runtime {
Some(message) => {
let text = runtime_message_text(&message)?;
(message, text)
}
None => {
break_messages::compile_background(
event_id,
&user,
&request_context,
&action_context,
blob_sync,
)
.await?
}
};
user.text = text;
turn_user = Some(user);
vec![message]
}
(Some(user), Some(event_id)) => {
let runtime = match existing_runtime {
Some(message) => message,
None => {
break_messages::compile(
event_id,
checkpoint_mode,
&user,
&request_context,
&action_context,
compiler,
blob_sync,
)
.await?
}
};
request_context_message
.into_iter()
.chain(std::iter::once(runtime))
.collect()
}
(None, None) => Vec::new(),
_ => {
return Err(Error::Protocol(
"Cursor action has an incomplete runtime event".into(),
))
}
}
};
let (base_checkpoint_id, reused) = store
.match_checkpoint_prefix(&conversation_id, base_checkpoint_id, &initial_messages)
.await?;
initial_messages.drain(..reused);
let action = if compacting {
RunAction::Compact
} else if starts_turn {
RunAction::Start
} else {
let pending_tool_round = match request
.conversation_state
.as_ref()
.map(|state| state.pending_tool_calls.as_slice())
.unwrap_or_default()
{
[] => None,
[pending] => Some(messages::decode_pending(pending)?),
pending => {
return Err(Error::Protocol(format!(
"Cursor resume contains {} pending assistant messages",
pending.len()
)))
}
};
RunAction::Resume { pending_tool_round }
};
let exec = exec_context(
request,
&request_context,
&conversation_id,
&model.model_id,
subagents_disabled,
&subagent_model_overrides,
);
Ok((
PreparedRun {
run_id,
cursor_request_id: Some(request_id.into()),
conversation_id,
kind: RunKind::Root,
model,
prompt,
initial_messages,
action,
base_checkpoint_id,
},
CursorRunContext {
request_id: request_id.into(),
mode: mode_number,
turn_user,
exec,
dynamic_tools: dynamic
.into_iter()
.map(|(name, (wire, _))| (name, wire))
.collect(),
checkpoint_prompt,
compacting,
background_completion,
},
))
}
fn runtime_message_text(message: &CanonicalMessage) -> Result<String> {
let MessageContent::Parts { parts } = &message.content else {
return Err(Error::Protocol(
"stored runtime message does not contain parts".into(),
));
};
let Some(ContentPart::Text { text }) = parts.first() else {
return Err(Error::Protocol(
"stored runtime message does not start with text".into(),
));
};
Ok(text.clone())
}
fn validate_prompt_root(messages: &[CanonicalMessage]) -> Result<()> {
let prompts = messages
.iter()
.filter(|message| message.role == Role::System && message.origin == Origin::Prompt)
.collect::<Vec<_>>();
let [prompt] = prompts.as_slice() else {
return Err(Error::Protocol(format!(
"Cursor history contains {} system prompt roots",
prompts.len()
)));
};
let MessageContent::Parts { parts } = &prompt.content else {
return Err(Error::Protocol(
"Cursor system prompt root is not textual content".into(),
));
};
let [ContentPart::Text { .. }] = parts.as_slice() else {
return Err(Error::Protocol(
"Cursor system prompt root is not one text part".into(),
));
};
Ok(())
}
fn execution_run_id(request_id: &str) -> RunId {
let execution_id = Uuid::new_v4().simple().to_string();
RunId::new(format!("{request_id}:{}", &execution_id[..8]))
}
fn action(request: &pb::AgentRunRequest) -> Result<ActionProjection> {
let conversation_mode = request
.conversation_state
.as_ref()
.and_then(|state| state.mode);
let mode = conversation_mode.unwrap_or(pb::AgentMode::Agent as i32);
let Some(action) = request
.action
.as_ref()
.and_then(|action| action.action.as_ref())
else {
return Ok(ActionProjection {
mode,
turn_user: None,
action_context: String::new(),
event_id: None,
input_id: None,
starts_turn: false,
compacting: false,
background_completion: false,
});
};
match action {
pb::conversation_action::Action::UserMessageAction(action) => {
let user = action.user_message.as_ref().ok_or_else(|| {
Error::Protocol("Cursor user message action has no UserMessage".into())
})?;
let mode = if user.mode == pb::AgentMode::Unspecified as i32 {
conversation_mode.unwrap_or(user.mode)
} else {
user.mode
};
if user.message_id.is_empty() {
return Err(Error::Protocol(
"Cursor user message action has no message_id".into(),
));
}
if user.text.trim() == "/summarize" {
return Ok(ActionProjection {
mode,
turn_user: Some(user.clone()),
action_context: String::new(),
event_id: None,
input_id: None,
starts_turn: false,
compacting: true,
background_completion: false,
});
}
let mut context = action
.prepend_user_messages
.iter()
.map(|message| message.text.trim())
.filter(|text| !text.is_empty())
.map(str::to_string)
.collect::<Vec<_>>();
context.extend(
user.subagent_system_reminder
.iter()
.filter(|text| !text.is_empty())
.cloned(),
);
let input_id = format!("cursor:user:{}", user.message_id);
Ok(ActionProjection {
mode,
turn_user: Some(user.clone()),
action_context: context.join("\n\n"),
event_id: None,
input_id: Some(input_id),
starts_turn: true,
compacting: false,
background_completion: false,
})
}
pb::conversation_action::Action::BackgroundTaskCompletionAction(action) => {
let projection = insert_messages::project(action, mode)?;
let event_id = projection.turn_user.message_id.clone();
Ok(ActionProjection {
mode,
action_context: projection.context,
event_id: Some(event_id),
input_id: None,
turn_user: Some(projection.turn_user),
starts_turn: true,
compacting: false,
background_completion: true,
})
}
pb::conversation_action::Action::ExecutePlanAction(action) => execute_plan(action),
pb::conversation_action::Action::SummarizeAction(_) => Ok(ActionProjection {
mode,
turn_user: None,
action_context: String::new(),
event_id: None,
input_id: None,
starts_turn: false,
compacting: true,
background_completion: false,
}),
_ => Ok(ActionProjection {
mode,
turn_user: None,
action_context: String::new(),
event_id: None,
input_id: None,
starts_turn: false,
compacting: false,
background_completion: false,
}),
}
}
fn execute_plan(action: &pb::ExecutePlanAction) -> Result<ActionProjection> {
let plan = action
.plan_file_content
.as_deref()
.or_else(|| action.plan.as_ref().map(|plan| plan.plan.as_str()))
.filter(|plan| !plan.trim().is_empty())
.ok_or_else(|| Error::Protocol("ExecutePlan is missing plan content".into()))?;
let source = action
.plan_file_uri
.as_deref()
.or(action.plan_file_path.as_deref())
.filter(|source| !source.is_empty());
let action_context = match source {
Some(source) => {
format!("<approved_plan>\n<plan_file>{source}</plan_file>\n{plan}\n</approved_plan>")
}
None => format!("<approved_plan>\n{plan}\n</approved_plan>"),
};
let identity = BlobId::digest(
format!(
"{}\0{}\0{}\0{}\0{}",
action.execution_mode,
action.plan_id.as_deref().unwrap_or_default(),
action.kickoff_message_id.as_deref().unwrap_or_default(),
source.unwrap_or_default(),
plan,
)
.as_bytes(),
)
.to_base64();
let event_id = format!("execute-plan:{identity}");
Ok(ActionProjection {
mode: action.execution_mode,
turn_user: Some(pb::UserMessage {
text: "Execute the approved plan.".into(),
message_id: event_id.clone(),
mode: action.execution_mode,
..Default::default()
}),
action_context,
event_id: Some(event_id),
input_id: None,
starts_turn: true,
compacting: false,
background_completion: false,
})
}
pub(super) fn mode_from_proto(mode: i32) -> Result<Mode> {
let mode = pb::AgentMode::try_from(mode)
.map_err(|_| Error::Protocol(format!("unknown Cursor agent mode: {mode}")))?;
match mode {
pb::AgentMode::Agent => Ok(Mode::Agent),
pb::AgentMode::Ask => Ok(Mode::Ask),
pb::AgentMode::Plan => Ok(Mode::Plan),
pb::AgentMode::Debug => Ok(Mode::Debug),
pb::AgentMode::Multitask => Ok(Mode::Multitask),
mode => Err(Error::Protocol(format!(
"unsupported Cursor agent mode: {}",
mode.as_str_name()
))),
}
}
fn exec_context(
request: &pb::AgentRunRequest,
request_context: &pb::RequestContext,
conversation_id: &ConversationId,
model_id: &str,
subagents_disabled: bool,
overrides: &[(
crate::model::SubagentKind,
crate::model::SubagentModelOverride,
)],
) -> ExecContext {
let subagent_model = overrides.first().map(|(_, value)| match value {
crate::model::SubagentModelOverride::Explicit(model) => {
SubagentModel::Model(model.model_id.clone())
}
crate::model::SubagentModelOverride::Inherit => SubagentModel::Model(model_id.into()),
crate::model::SubagentModelOverride::Disabled => SubagentModel::Disabled,
});
ExecContext {
conversation_id: conversation_id.to_string(),
root_conversation_id: request
.conversation_group_id
.clone()
.unwrap_or_else(|| conversation_id.to_string()),
default_subagent_model: model_id.into(),
subagent_model,
allow_subagents: request.subagent_type_name.is_none() && !subagents_disabled,
subagents_disabled,
terminals_folder: request_context
.env
.as_ref()
.map(|env| env.terminals_folder.clone())
.unwrap_or_default(),
admin_command_denylist: request_context.admin_command_denylist.clone(),
mcp_routes: context::meta_mcp_routes(request_context),
}
}
+13
View File
@@ -0,0 +1,13 @@
//! Defines commands accepted by a Conversation runtime.
use crate::cursor::protocol::proto::agent::v1 as pb;
#[derive(Debug)]
pub enum TransportCommand {
Append {
seqno: i64,
message: Box<pb::AgentClientMessage>,
},
Disconnect,
Close,
}
@@ -0,0 +1,11 @@
//! Applies Ignore, InsertMessages, and BreakMessages delivery semantics.
use crate::{model::RunId, run::CommandResult};
pub use crate::cursor::compile::{CompiledMessages, MessageDelivery};
pub fn target_result(target: Option<&RunId>, current: &RunId) -> Option<CommandResult> {
target
.filter(|target| *target != current)
.map(|_| CommandResult::StaleTarget)
}
+15
View File
@@ -0,0 +1,15 @@
//! Owns conversation-scoped runtime coordination.
mod command;
mod delivery;
mod output;
mod pending;
mod registry;
mod runtime;
pub use command::*;
pub use delivery::*;
pub(crate) use output::*;
pub(crate) use pending::*;
pub use registry::*;
pub(crate) use runtime::*;
File diff suppressed because it is too large Load Diff
+26
View File
@@ -0,0 +1,26 @@
//! Stores messages waiting across Run lifecycle boundaries.
use std::collections::{HashSet, VecDeque};
use super::CompiledMessages;
#[derive(Default)]
pub struct PendingMessages {
queued: VecDeque<CompiledMessages>,
event_ids: HashSet<String>,
}
impl PendingMessages {
pub fn push(&mut self, messages: CompiledMessages) -> bool {
if !self.event_ids.insert(messages.event_id.clone()) {
return false;
}
self.queued.push_back(messages);
true
}
pub fn drain(&mut self) -> impl Iterator<Item = CompiledMessages> + '_ {
self.event_ids.clear();
self.queued.drain(..)
}
}
+196
View File
@@ -0,0 +1,196 @@
//! Maps conversation IDs to active conversation runtimes.
use std::{collections::HashMap, sync::Arc};
use tokio::sync::{mpsc, Mutex, Notify};
use crate::{
cursor::{prompting::PromptCompiler, transport::TransportHandle},
model::{ConversationId, RunId},
provider::Provider,
run::{CommandResult, RunHandle},
store::Store,
};
use super::{CompiledMessages, MessageDelivery, PendingMessages, TransportCommand};
#[derive(Clone)]
pub struct ConversationRegistry {
inner: Arc<RegistryInner>,
}
#[derive(Clone)]
pub(crate) struct ConversationDependencies {
pub store: Store,
pub provider: Arc<dyn Provider>,
pub compiler: PromptCompiler,
}
struct RegistryInner {
current: Mutex<HashMap<ConversationId, ActiveRun>>,
pending: Mutex<HashMap<ConversationId, PendingMessages>>,
changed: Notify,
pub dependencies: ConversationDependencies,
}
#[derive(Clone)]
struct ActiveRun {
run_id: RunId,
handle: RunHandle,
}
impl ConversationRegistry {
pub fn new(store: Store, provider: Arc<dyn Provider>, compiler: PromptCompiler) -> Self {
Self {
inner: Arc::new(RegistryInner {
current: Mutex::new(HashMap::new()),
pending: Mutex::new(HashMap::new()),
changed: Notify::new(),
dependencies: ConversationDependencies {
store,
provider,
compiler,
},
}),
}
}
pub(crate) fn dependencies(&self) -> &ConversationDependencies {
&self.inner.dependencies
}
pub(crate) fn bind_transport(
&self,
handle: TransportHandle,
receiver: mpsc::Receiver<TransportCommand>,
) {
super::ConversationRuntime::spawn(self.clone(), handle, receiver);
}
pub(crate) async fn activate(
&self,
conversation_id: ConversationId,
run_id: RunId,
handle: RunHandle,
) {
let previous = self.inner.current.lock().await.insert(
conversation_id,
ActiveRun {
run_id: run_id.clone(),
handle,
},
);
if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) {
previous.handle.cancel();
}
}
pub async fn deliver(
&self,
conversation_id: &ConversationId,
compiled: CompiledMessages,
) -> CommandResult {
if compiled.delivery == MessageDelivery::Ignore {
return CommandResult::Applied;
}
let active = self
.inner
.current
.lock()
.await
.get(conversation_id)
.cloned();
let Some(active) = active else {
self.inner
.pending
.lock()
.await
.entry(conversation_id.clone())
.or_default()
.push(compiled);
return CommandResult::RunEnded;
};
if compiled
.target_run_id
.as_ref()
.is_some_and(|target| target != &active.run_id)
{
return CommandResult::StaleTarget;
}
let pending = compiled.clone();
let result = match compiled.delivery {
MessageDelivery::Ignore => CommandResult::Applied,
MessageDelivery::InsertMessages => {
active
.handle
.insert_messages(compiled.event_id, compiled.messages)
.await
}
MessageDelivery::BreakMessages => {
active
.handle
.break_messages(compiled.event_id, compiled.messages)
.await
}
};
if matches!(result, CommandResult::RunClosing | CommandResult::RunEnded) {
self.inner
.pending
.lock()
.await
.entry(conversation_id.clone())
.or_default()
.push(pending);
}
result
}
pub(crate) async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) {
let mut current = self.inner.current.lock().await;
if current
.get(conversation_id)
.is_some_and(|run| &run.run_id == run_id)
{
current.remove(conversation_id);
self.inner.changed.notify_waiters();
}
}
pub(crate) async fn wait_until_idle(&self, conversation_id: &ConversationId) {
loop {
let changed = self.inner.changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
if !self
.inner
.current
.lock()
.await
.contains_key(conversation_id)
{
return;
}
changed.await;
}
}
pub(crate) async fn take_pending(
&self,
conversation_id: &ConversationId,
) -> Vec<CompiledMessages> {
self.inner
.pending
.lock()
.await
.remove(conversation_id)
.map(|mut pending| pending.drain().collect())
.unwrap_or_default()
}
pub async fn shutdown(&self) {
let current = std::mem::take(&mut *self.inner.current.lock().await);
for active in current.into_values() {
active.handle.cancel();
}
}
}
+651
View File
@@ -0,0 +1,651 @@
//! Owns the current Run and coordinates the Conversation lifecycle.
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::{
cursor::{
checkpoint::CheckpointBuilder,
compile,
protocol::proto::agent::v1 as pb,
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
tools::{
codec,
runtime::CursorToolRuntime,
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
ClientToolEvent, ToolDispatcher,
},
transport::{OrderedInbox, TransportHandle},
},
run::{CommandResult, RunEngine, RunHandle},
};
use super::{
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
ConversationRegistry, MessageDelivery, TransportCommand,
};
pub struct ConversationRuntime;
#[derive(Clone)]
struct RunGeneration {
superseded: CancellationToken,
finished: CancellationToken,
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
results: ToolResultSender,
runtime_actions: mpsc::UnboundedSender<compile::RuntimeAction>,
tool_runtime: CursorToolRuntime,
tools: ToolDispatcher,
}
struct FinishGeneration(CancellationToken);
impl Drop for FinishGeneration {
fn drop(&mut self) {
self.0.cancel();
}
}
impl ConversationRuntime {
pub(crate) fn spawn(
registry: ConversationRegistry,
handle: TransportHandle,
mut receiver: mpsc::Receiver<TransportCommand>,
) {
tokio::spawn(async move {
let dependencies = registry.dependencies().clone();
let blob_sync = BlobSynchronizer::new(
handle.request_id().into(),
dependencies.store.clone(),
handle.clone(),
);
let mut inbox = OrderedInbox::starting_at(0);
let tool_runtime_factory = CursorToolRuntime::default();
let context_sync =
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
let mut current = None::<RunGeneration>;
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();
}
}
super::finish_cancelled(&handle).ok();
break;
}
};
match command {
TransportCommand::Disconnect => {
handle.mark_disconnected();
if let Some(generation) = current.as_ref() {
generation.superseded.cancel();
if let Some(run) = generation.run.lock().clone() {
run.cancel();
}
for id in generation.tool_runtime.drain_running().await {
let _ = handle.emit(&codec::abort(id));
}
}
super::finish_cancelled(&handle).ok();
break;
}
TransportCommand::Close => {
break;
}
TransportCommand::Append { seqno, message } => {
for (_seqno, message) in inbox.push(seqno, *message) {
{
match message.message {
Some(pb::agent_client_message::Message::RunRequest(
request,
)) => {
if let Some(conversation_id) =
request.conversation_id.as_deref()
{
if let Err(error) =
handle.set_conversation_id(conversation_id)
{
tracing::error!(
request_id = handle.request_id(),
%error,
"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(),
);
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(),
request,
dependencies.clone(),
blob_sync.clone(),
context_sync.clone(),
generation,
previous_finished,
result_receiver,
runtime_action_receiver,
);
}
Some(pb::agent_client_message::Message::ExecClientMessage(
message,
)) => {
if context_sync.handle_client(&message).await {
continue;
}
let Some(generation) = current.as_ref() else {
continue;
};
match codec::client_event(
&message,
&generation.tool_runtime,
)
.await
{
Ok(codec::ClientExecEvent::Delta(message)) => {
let _ = handle.emit(&message);
}
Ok(codec::ClientExecEvent::Message(message)) => {
let _ = handle.emit(&message);
}
Ok(codec::ClientExecEvent::Completed(result)) => {
generation.results.send(*result)
}
Ok(codec::ClientExecEvent::Pending) => {}
Err(error) => generation.results.send_error(error),
}
}
Some(
pb::agent_client_message::Message::ExecClientControlMessage(
message,
),
) => {
use pb::exec_client_control_message::Message;
match message.message {
Some(Message::StreamClose(close)) => {
if context_sync.handle_stream_close(close.id).await
{
continue;
}
let Some(generation) = current.as_ref() else {
continue;
};
match codec::stream_closed(
close.id,
&generation.tool_runtime,
)
.await
{
Ok(Some(completion)) => {
generation.results.send(completion)
}
Ok(None) => {}
Err(error) => {
generation.results.send_error(error)
}
}
}
Some(Message::Throw(throw)) => {
if context_sync
.handle_throw(
throw.id,
format!(
"Cursor request context failed: {}",
throw.error
),
)
.await
{
continue;
}
let Some(generation) = current.as_ref() else {
continue;
};
if generation
.tool_runtime
.is_interrupted(throw.id)
.await
{
generation
.tool_runtime
.discard_exec(throw.id)
.await;
continue;
}
match generation
.tool_runtime
.take_exec(throw.id)
.await
{
Some(pending) => generation.results.send_error(
crate::Error::Protocol(format!(
"Exec {} failed: {}",
pending.call.call_id, throw.error
)),
),
None => generation.results.send_error(
crate::Error::Protocol(format!(
"unknown ExecClientThrow id: {}",
throw.id
)),
),
}
}
Some(Message::Heartbeat(_)) | None => {}
}
}
Some(
pb::agent_client_message::Message::InteractionResponse(
message,
),
) => {
let Some(generation) = current.as_ref() else {
continue;
};
match generation.tools.interaction_response(&message).await
{
Ok(ClientToolEvent::Completed(completion)) => {
generation.results.send(*completion)
}
Ok(ClientToolEvent::Pending) => {}
Err(error) => generation.results.send_error(error),
}
}
Some(pb::agent_client_message::Message::KvClientMessage(
message,
)) => {
let _ = blob_sync.handle_client(message).await;
}
// TODO: ConversationAction has two different delivery paths that
// must not be conflated:
//
// 1. AgentRunRequest.action starts/resumes a Run. compile::prepare
// currently consumes UserMessageAction,
// BackgroundTaskCompletionAction, SummarizeAction and
// ExecutePlanAction. ResumeAction only works indirectly through
// the absence of a new runtime event and still needs an explicit
// implementation that consumes ResumeAction.request_context.
// 2. AgentClientMessage::ConversationAction arrives while a Bidi Run
// is already active and needs a runtime dispatcher here. Supporting
// an Action in compile::prepare does not mean this path supports it.
//
// Cursor 3.16 sends a queued follow-up as InjectContextAction.
// It targets expected_run_id and asks the active Run to yield to the
// queued message. The session owns this path because interruption must
// abort active execs and publish a recoverable checkpoint before the
// old Run ends. It must not be reduced to handle.cancel() here.
//
// The remaining unimplemented Action variants are
// ShellCommandAction, StartPlanAction,
// AsyncAskQuestionCompletionAction, BackgroundShellAction,
// BackgroundSubagentAction,
// SubscriptionNotificationAction and GoalContinuationAction.
// Variants whose wire behavior is not captured yet need evidence
// before assigning semantics. Every unsupported runtime Action must
// return an explicit Protocol Error rather than falling through silently.
Some(
pb::agent_client_message::Message::ConversationAction(
action,
),
) => match action.action {
Some(
pb::conversation_action::Action::UserMessageAction(
action,
),
) => {
let Some(generation) = 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(),
));
}
}
Some(pb::conversation_action::Action::CancelAction(_)) => {
if let Some(generation) = current.as_ref() {
if let Some(run) = generation.run.lock().clone() {
run.cancel();
}
for id in
generation.tool_runtime.drain_running().await
{
let _ = handle.emit(&codec::abort(id));
}
}
}
Some(
pb::conversation_action::Action::InjectContextAction(
action,
),
) => {
let Some(generation) = current.as_ref() else {
continue;
};
if generation
.runtime_actions
.send(compile::RuntimeAction::Inject(action))
.is_err()
{
generation.results.send_error(crate::Error::Protocol(
"InjectContextAction arrived without an active Run"
.into(),
));
}
}
Some(
pb::conversation_action::Action::CancelSubagentAction(
action,
),
) => {
if let Some(generation) = current.as_ref() {
if let Some(id) = generation
.tool_runtime
.running_task_exec_id(&action.subagent_id)
.await
{
let _ = handle.emit(&codec::abort(id));
}
}
}
Some(action) => {
tracing::warn!(
request_id = handle.request_id(),
action = runtime_action_name(&action),
"ignoring unsupported runtime ConversationAction"
);
}
None => {
if let Some(generation) = current.as_ref() {
generation.results.send_error(
crate::Error::Protocol(
"runtime ConversationAction has no action"
.into(),
),
);
}
}
},
_ => {}
}
}
}
}
}
}
});
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_run_request(
registry: ConversationRegistry,
handle: TransportHandle,
request: pb::AgentRunRequest,
dependencies: ConversationDependencies,
blob_sync: BlobSynchronizer,
context_sync: RequestContextSynchronizer,
generation: RunGeneration,
previous_finished: Option<CancellationToken>,
results: ToolResultReceiver,
runtime_actions: mpsc::UnboundedReceiver<compile::RuntimeAction>,
) {
tokio::spawn(async move {
let _finished = FinishGeneration(generation.finished.clone());
if let Some(previous_finished) = previous_finished {
tokio::select! {
biased;
_ = generation.superseded.cancelled() => return,
_ = previous_finished.cancelled() => {}
}
}
if generation.superseded.is_cancelled() {
return;
}
let mut checkpoint = CheckpointBuilder::new(
dependencies.store.clone(),
blob_sync.clone(),
handle.parent().map(|parent| parent.tool_call_id.clone()),
request.conversation_state.clone(),
);
let prepared = tokio::select! {
biased;
_ = generation.superseded.cancelled() => return,
prepared = compile::prepare(
handle.request_id(),
&request,
compile::PrepareDependencies {
compiler: &dependencies.compiler,
store: &dependencies.store,
checkpoint: &checkpoint,
blob_sync: &blob_sync,
context_sync: &context_sync,
},
) => prepared,
};
let (mut prepared, context) = match prepared {
Ok(prepared) => prepared,
Err(error) => {
if generation.superseded.is_cancelled() {
return;
}
tracing::error!(
request_id = handle.request_id(),
%error,
"failed to prepare Cursor Run"
);
let _ = super::finish_failed(&handle, &error);
let _ = handle.command(TransportCommand::Close).await;
return;
}
};
checkpoint.configure(
prepared.model.model_id.clone(),
prepared.model.context_window_tokens,
context.checkpoint_prompt.instructions.clone(),
context.checkpoint_prompt.tools.clone(),
context.dynamic_tools.keys().cloned().collect(),
context.turn_user.clone(),
);
if context.background_completion {
let event_id = prepared
.initial_messages
.first()
.and_then(|message| message.runtime_event_id.clone())
.unwrap_or_else(|| format!("background:{}", prepared.run_id));
match registry
.deliver(
&prepared.conversation_id,
CompiledMessages {
event_id,
target_run_id: None,
messages: prepared.initial_messages.clone(),
delivery: MessageDelivery::InsertMessages,
},
)
.await
{
CommandResult::Applied | CommandResult::Duplicate => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
}
return;
}
CommandResult::RunClosing => {
prepared.initial_messages.clear();
tokio::select! {
biased;
_ = generation.superseded.cancelled() => return,
_ = registry.wait_until_idle(&prepared.conversation_id) => {}
}
if let Ok(checkpoint) = dependencies
.store
.ensure_conversation(&prepared.conversation_id)
.await
{
prepared.base_checkpoint_id = checkpoint;
}
}
CommandResult::RunEnded => {
prepared.initial_messages.clear();
}
CommandResult::StaleTarget => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
}
return;
}
}
}
let pending = registry.take_pending(&prepared.conversation_id).await;
if !pending.is_empty() {
let mut messages = pending
.into_iter()
.flat_map(|pending| pending.messages)
.collect::<Vec<_>>();
messages.extend(prepared.initial_messages);
prepared.initial_messages = messages;
if let Ok(checkpoint) = dependencies
.store
.ensure_conversation(&prepared.conversation_id)
.await
{
prepared.base_checkpoint_id = checkpoint;
}
}
if generation.superseded.is_cancelled() {
return;
}
let run_id = prepared.run_id.clone();
let conversation_id = prepared.conversation_id.clone();
let (port, core, run_handle) = crate::run::channel(run_id.clone(), 256);
*generation.run.lock() = Some(run_handle.clone());
if generation.superseded.is_cancelled() {
run_handle.cancel();
*generation.run.lock() = None;
return;
}
registry
.activate(conversation_id.clone(), run_id.clone(), run_handle.clone())
.await;
let cancellation = run_handle.cancellation();
let engine = RunEngine::new(dependencies.store.clone(), dependencies.provider.clone());
let core_run = tokio::spawn(async move { engine.run(prepared, port, cancellation).await });
let output = ConversationOutput::new(
handle.clone(),
dependencies.store.clone(),
context,
core,
run_handle,
registry.clone(),
ConversationOutputDependencies {
superseded: generation.superseded.clone(),
tools: generation.tools.clone(),
results,
runtime_actions,
compiler: dependencies.compiler.clone(),
blob_sync,
checkpoint,
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 _ = core_run.await;
registry.release(&conversation_id, &run_id).await;
if generation
.run
.lock()
.as_ref()
.is_some_and(|run| run.run_id() == &run_id)
{
*generation.run.lock() = None;
}
if !generation.superseded.is_cancelled() {
let _ = handle.command(TransportCommand::Close).await;
}
});
}
fn runtime_action_name(action: &pb::conversation_action::Action) -> &'static str {
use pb::conversation_action::Action;
match action {
Action::UserMessageAction(_) => "UserMessageAction",
Action::ResumeAction(_) => "ResumeAction",
Action::CancelAction(_) => "CancelAction",
Action::SummarizeAction(_) => "SummarizeAction",
Action::ShellCommandAction(_) => "ShellCommandAction",
Action::StartPlanAction(_) => "StartPlanAction",
Action::ExecutePlanAction(_) => "ExecutePlanAction",
Action::AsyncAskQuestionCompletionAction(_) => "AsyncAskQuestionCompletionAction",
Action::CancelSubagentAction(_) => "CancelSubagentAction",
Action::BackgroundTaskCompletionAction(_) => "BackgroundTaskCompletionAction",
Action::BackgroundShellAction(_) => "BackgroundShellAction",
Action::BackgroundSubagentAction(_) => "BackgroundSubagentAction",
Action::SubscriptionNotificationAction(_) => "SubscriptionNotificationAction",
Action::GoalContinuationAction(_) => "GoalContinuationAction",
Action::InjectContextAction(_) => "InjectContextAction",
}
}
+13
View File
@@ -0,0 +1,13 @@
//! Exposes the Cursor protocol adapter and its conversation runtime.
pub mod checkpoint;
pub mod compile;
pub mod conversation;
pub mod prompting;
pub mod protocol;
pub mod services;
pub mod tools;
pub mod transport;
pub use conversation::TransportCommand;
pub use transport::{TransportHandle, TransportParent, TransportRegistry, TransportRoute};
+180
View File
@@ -0,0 +1,180 @@
//! Loads embedded Cursor Prompt assets.
use std::{path::Path, sync::OnceLock};
use crate::{model::ToolDefinition, Error, Result};
use super::catalog::Catalog;
static EMBEDDED_PROMPTS: include_dir::Dir<'_> =
include_dir::include_dir!("$CARGO_MANIFEST_DIR/prompt/cursor");
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Mode {
Agent,
Ask,
Plan,
Debug,
Multitask,
Subagent,
Compaction,
}
impl Mode {
pub fn parse(value: &str) -> Result<Self> {
match value.to_ascii_lowercase().as_str() {
"agent" => Ok(Self::Agent),
"ask" => Ok(Self::Ask),
"plan" => Ok(Self::Plan),
"debug" => Ok(Self::Debug),
"multitask" => Ok(Self::Multitask),
"subagent" => Ok(Self::Subagent),
"compaction" => Ok(Self::Compaction),
other => Err(Error::Config(format!("unknown prompt mode: {other}"))),
}
}
fn name(self) -> &'static str {
match self {
Self::Agent => "agent",
Self::Ask => "ask",
Self::Plan => "plan",
Self::Debug => "debug",
Self::Multitask => "multitask",
Self::Subagent => "subagent",
Self::Compaction => "compaction",
}
}
fn index(self) -> usize {
match self {
Self::Agent => 0,
Self::Ask => 1,
Self::Plan => 2,
Self::Debug => 3,
Self::Multitask => 4,
Self::Subagent => 5,
Self::Compaction => 6,
}
}
}
#[derive(Clone, Debug)]
pub struct ModeAssets {
pub prompt: String,
pub runtime: String,
pub tools: Vec<ToolDefinition>,
}
#[derive(Clone, Debug)]
pub struct PromptAssets {
modes: [ModeAssets; 7],
}
impl PromptAssets {
pub fn load(root: &Path) -> Result<Self> {
Self::read(|path| {
let path = root.join(path);
path.exists()
.then(|| std::fs::read_to_string(path).map_err(Error::from))
.transpose()
})
}
pub fn embedded() -> Result<Self> {
Self::read(|path| {
EMBEDDED_PROMPTS
.get_file(path)
.map(|file| {
file.contents_utf8()
.map(str::to_string)
.ok_or_else(|| Error::Config(format!("prompt asset is not UTF-8: {path}")))
})
.transpose()
})
}
fn read(mut asset: impl FnMut(&str) -> Result<Option<String>>) -> Result<Self> {
let catalog = Catalog::parse(
&asset("tools.json")?
.ok_or_else(|| Error::Config("missing Cursor tools.json".into()))?,
)?;
let mut modes = Vec::with_capacity(7);
for mode in [
Mode::Agent,
Mode::Ask,
Mode::Plan,
Mode::Debug,
Mode::Multitask,
Mode::Subagent,
Mode::Compaction,
] {
let prompt = asset(&format!("{}/prompt.md", mode.name()))?
.ok_or_else(|| Error::Config(format!("missing prompt for {mode:?}")))?;
let runtime = asset(&format!("{}/runtime.md", mode.name()))?
.ok_or_else(|| Error::Config(format!("missing runtime template for {mode:?}")))?;
validate_runtime_template(mode, &runtime)?;
let manifest = asset(&format!("modes/{}.json", mode.name()))?
.ok_or_else(|| Error::Config(format!("missing manifest for {mode:?}")))?;
let tools = catalog.select_json(&manifest)?;
modes.push(ModeAssets {
prompt,
runtime,
tools,
});
}
Ok(Self {
modes: modes
.try_into()
.map_err(|_| Error::Config("incomplete Cursor prompt mode catalog".into()))?,
})
}
pub fn mode(&self, mode: Mode) -> &ModeAssets {
&self.modes[mode.index()]
}
}
const RUNTIME_VARIABLES: &[&str] = &[
"OPEN_FILES",
"SELECTED_CONTEXT",
"ACTION_CONTEXT",
"TIMESTAMP",
"USER_QUERY",
"DEBUG_SERVER_ENDPOINT",
"DEBUG_LOG_PATH",
"DEBUG_SESSION_ID",
];
fn validate_runtime_template(mode: Mode, template: &str) -> Result<()> {
let expression = runtime_expression();
for capture in expression.captures_iter(template) {
let name = &capture[1];
if !RUNTIME_VARIABLES.contains(&name) {
return Err(Error::Config(format!(
"unknown variable in {mode:?} runtime template: {name}"
)));
}
}
for required in ["TIMESTAMP", "USER_QUERY"] {
let token = format!("{{{{{required}}}}}");
if !template.contains(&token) {
return Err(Error::Config(format!(
"{mode:?} runtime template is missing {token}"
)));
}
}
let stripped = expression.replace_all(template, "");
if stripped.contains("{{") || stripped.contains("}}") {
return Err(Error::Config(format!(
"malformed placeholder in {mode:?} runtime template"
)));
}
Ok(())
}
pub(super) fn runtime_expression() -> &'static regex::Regex {
static EXPRESSION: OnceLock<regex::Regex> = OnceLock::new();
EXPRESSION.get_or_init(|| {
regex::Regex::new(r"\{\{([A-Z_]+)\}\}").expect("valid runtime placeholder expression")
})
}
+3
View File
@@ -0,0 +1,3 @@
//! Routes Prompt asset loading through the Tool schema registry.
pub(super) use crate::cursor::tools::registry::ToolRegistry as Catalog;
+94
View File
@@ -0,0 +1,94 @@
//! Compiles stable Prompt specifications for Cursor modes.
use std::collections::BTreeMap;
use crate::{
model::{ModelSpec, PromptSpec, ToolDefinition},
Error, Result,
};
use super::{assets::runtime_expression, Mode, PromptAssets};
#[derive(Clone)]
pub struct PromptCompiler {
assets: PromptAssets,
}
impl PromptCompiler {
pub fn new(assets: PromptAssets) -> Self {
Self { assets }
}
pub fn runtime_message(&self, mode: Mode, values: &BTreeMap<&str, String>) -> Result<String> {
render(&self.assets.mode(mode).runtime, values)
}
pub fn prompt_spec(
&self,
mode: Mode,
model: &ModelSpec,
dynamic_tools: &[ToolDefinition],
suppress_subagent_progress: bool,
) -> Result<PromptSpec> {
let mut tools = self.tools(mode, suppress_subagent_progress);
let mut dynamic_tools = dynamic_tools.to_vec();
dynamic_tools.sort_by(|left, right| left.name.cmp(&right.name));
append_dynamic_tools(&mut tools, dynamic_tools)?;
if !model.supports_image_generation {
tools.retain(|tool| tool.name != "GenerateImage");
}
let fake_model_name = model
.display_name
.as_deref()
.unwrap_or(model.model_id.as_str());
Ok(PromptSpec {
instructions: self
.assets
.mode(mode)
.prompt
.replace("{{FAKE_MODEL_NAME}}", fake_model_name),
tools,
})
}
fn tools(&self, mode: Mode, suppress_subagent_progress: bool) -> Vec<ToolDefinition> {
let mut tools = self.assets.mode(mode).tools.clone();
if mode == Mode::Subagent && suppress_subagent_progress {
tools.retain(|tool| tool.name != "UpdateCurrentStep");
}
tools
}
}
fn render(template: &str, values: &BTreeMap<&str, String>) -> Result<String> {
let expression = runtime_expression();
let mut output = String::with_capacity(template.len());
let mut cursor = 0;
for capture in expression.captures_iter(template) {
let token = capture.get(0).expect("runtime template token");
let name = &capture[1];
let value = values
.get(name)
.ok_or_else(|| Error::Protocol(format!("runtime template value is missing: {name}")))?;
output.push_str(&template[cursor..token.start()]);
output.push_str(value);
cursor = token.end();
}
output.push_str(&template[cursor..]);
Ok(output.trim().to_string())
}
fn append_dynamic_tools(
tools: &mut Vec<ToolDefinition>,
additions: Vec<ToolDefinition>,
) -> Result<()> {
for tool in additions {
if tools.iter().any(|existing| existing.name == tool.name) {
return Err(Error::Protocol(format!(
"dynamic MCP tool conflicts with a mode tool: {}",
tool.name
)));
}
tools.push(tool);
}
Ok(())
}
@@ -0,0 +1,83 @@
//! Builds deterministic prompt state derived from Conversation context.
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::model::{CanonicalMessage, MessageContent};
#[derive(Clone, Debug, Default, Serialize, Deserialize, PartialEq)]
pub struct DerivedState {
pub todos: Option<Value>,
pub plan: Option<Value>,
}
pub fn fold_derived_state(messages: &[CanonicalMessage]) -> DerivedState {
let mut state = DerivedState::default();
let mut calls = std::collections::HashMap::<String, (String, Value)>::new();
for message in messages {
match &message.content {
MessageContent::Assistant { tool_calls, .. } => {
for call in tool_calls {
calls.insert(
call.call_id.clone(),
(call.name.clone(), call.arguments.clone()),
);
}
}
MessageContent::ToolResult(result) if !result.is_error => {
let Some((name, input)) = calls.get(&result.call_id).cloned() else {
continue;
};
match normalize(&name).as_str() {
"todowrite" | "updatetodos" => {
state.todos = Some(apply_todo_write(state.todos.take(), input));
}
"createplan" | "updateplan" | "writeplan" => state.plan = Some(input),
_ => {}
}
}
_ => {}
}
}
state
}
fn apply_todo_write(current: Option<Value>, mut input: Value) -> Value {
if !input.get("merge").and_then(Value::as_bool).unwrap_or(false) {
return input;
}
let mut todos = current
.as_ref()
.and_then(|value| value.get("todos"))
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
let patches = input
.get("todos")
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
for patch in patches {
let existing = patch.get("id").and_then(Value::as_str).and_then(|id| {
todos
.iter_mut()
.find(|todo| todo.get("id").and_then(Value::as_str) == Some(id))
});
match (existing, patch) {
(Some(Value::Object(todo)), Value::Object(patch)) => todo.extend(patch),
(_, patch) => todos.push(patch),
}
}
if let Some(object) = input.as_object_mut() {
object.insert("merge".into(), Value::Bool(false));
object.insert("todos".into(), Value::Array(todos));
}
input
}
fn normalize(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
+10
View File
@@ -0,0 +1,10 @@
//! Exposes Cursor Prompt compilation.
mod assets;
mod catalog;
mod compiler;
mod derived_state;
pub use assets::*;
pub use compiler::*;
pub use derived_state::*;
+122
View File
@@ -0,0 +1,122 @@
//! Encodes and decodes Connect protocol frames.
use bytes::{BufMut, Bytes, BytesMut};
use prost::Message;
use serde::Serialize;
use crate::{Error, Result};
pub const END_STREAM_FLAG: u8 = 0x02;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ConnectCode {
Canceled,
InvalidArgument,
NotFound,
Unavailable,
Internal,
}
impl ConnectCode {
fn as_str(self) -> &'static str {
match self {
Self::Canceled => "canceled",
Self::InvalidArgument => "invalid_argument",
Self::NotFound => "not_found",
Self::Unavailable => "unavailable",
Self::Internal => "internal",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
pub struct ConnectErrorDetail {
#[serde(rename = "type")]
pub type_name: String,
pub value: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConnectStreamError {
pub code: ConnectCode,
pub message: String,
pub details: Vec<ConnectErrorDetail>,
}
#[derive(Serialize)]
struct EndStreamResponse<'a> {
error: WireError<'a>,
}
#[derive(Serialize)]
struct WireError<'a> {
code: &'static str,
#[serde(skip_serializing_if = "str::is_empty")]
message: &'a str,
#[serde(skip_serializing_if = "details_are_empty")]
details: &'a [ConnectErrorDetail],
}
fn details_are_empty(details: &&[ConnectErrorDetail]) -> bool {
details.is_empty()
}
pub fn encode_message<M: Message>(message: &M) -> Result<Bytes> {
let len = message.encoded_len();
let mut output = BytesMut::with_capacity(5 + len);
output.put_u8(0);
output.put_u32(len as u32);
message.encode(&mut output)?;
Ok(output.freeze())
}
pub fn encode_end_stream() -> Bytes {
encode_end_stream_payload(b"{}")
}
pub fn encode_error_end_stream(error: &ConnectStreamError) -> Result<Bytes> {
let payload = serde_json::to_vec(&EndStreamResponse {
error: WireError {
code: error.code.as_str(),
message: &error.message,
details: &error.details,
},
})?;
Ok(encode_end_stream_payload(&payload))
}
fn encode_end_stream_payload(payload: &[u8]) -> Bytes {
let mut output = BytesMut::with_capacity(5 + payload.len());
output.put_u8(END_STREAM_FLAG);
output.put_u32(payload.len() as u32);
output.extend_from_slice(payload);
output.freeze()
}
pub fn decode_unary<M: Message + Default>(body: &[u8]) -> Result<M> {
if body.len() >= 5 {
let flags = body[0];
let length = u32::from_be_bytes([body[1], body[2], body[3], body[4]]) as usize;
if flags & END_STREAM_FLAG == 0 && length == body.len() - 5 {
return Ok(M::decode(&body[5..])?);
}
}
Ok(M::decode(body)?)
}
pub fn decode_frames(mut body: &[u8]) -> Result<Vec<(u8, Bytes)>> {
let mut frames = Vec::new();
while !body.is_empty() {
if body.len() < 5 {
return Err(Error::Protocol("truncated Connect envelope".into()));
}
let flags = body[0];
let length = u32::from_be_bytes([body[1], body[2], body[3], body[4]]) as usize;
body = &body[5..];
if body.len() < length {
return Err(Error::Protocol("truncated Connect payload".into()));
}
frames.push((flags, Bytes::copy_from_slice(&body[..length])));
body = &body[length..];
}
Ok(frames)
}
+172
View File
@@ -0,0 +1,172 @@
//! Converts runtime events into live Cursor server messages.
use std::{collections::BTreeMap, time::Duration};
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec},
model::Usage,
provider::ModelEvent,
Result,
};
pub fn response_event(
event: &ModelEvent,
model_call_id: &str,
dynamic_mcp: &BTreeMap<String, pb::McpToolDefinition>,
) -> Result<Option<pb::AgentServerMessage>> {
use pb::interaction_update::Message;
let message = match event {
ModelEvent::TextDelta(text) => Message::TextDelta(pb::TextDeltaUpdate {
text: text.clone(),
is_server_notice: false,
}),
ModelEvent::ThinkingDelta(text) => Message::ThinkingDelta(pb::ThinkingDeltaUpdate {
text: text.clone(),
thinking_style: Some(pb::ThinkingStyle::Default as i32),
}),
ModelEvent::ToolCallStart { call_id, name, .. } => {
Message::PartialToolCall(pb::PartialToolCallUpdate {
call_id: call_id.clone(),
tool_call: Some(match dynamic_mcp.get(name) {
Some(definition) => codec::dynamic_mcp_placeholder(definition, call_id),
None => codec::tool_placeholder(name, call_id)?,
}),
args_text_delta: String::new(),
model_call_id: model_call_id.into(),
})
}
ModelEvent::ToolCallArgumentsDelta { .. } => return Ok(None),
ModelEvent::ToolCallEnd { .. }
| ModelEvent::Start { .. }
| ModelEvent::TextStart
| ModelEvent::TextEnd
| ModelEvent::ThinkingStart
| ModelEvent::ThinkingEnd
| ModelEvent::ProviderReplayState(_)
| ModelEvent::Usage(_)
| ModelEvent::Done(_) => return Ok(None),
};
Ok(Some(server_interaction(message)))
}
pub fn thinking_completed(elapsed: Duration) -> pb::AgentServerMessage {
let milliseconds = elapsed.as_millis().clamp(1, i32::MAX as u128) as i32;
server_interaction(pb::interaction_update::Message::ThinkingCompleted(
pb::ThinkingCompletedUpdate {
thinking_duration_ms: milliseconds,
},
))
}
pub fn heartbeat() -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::Heartbeat(
pb::HeartbeatUpdate {},
))
}
pub fn turn_ended(usage: Option<Usage>) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::TurnEnded(
pb::TurnEndedUpdate {
input_tokens: usage.and_then(|usage| usage.input_tokens.map(|value| value as i64)),
output_tokens: usage.and_then(|usage| usage.output_tokens.map(|value| value as i64)),
cache_read_tokens: usage
.and_then(|usage| usage.cache_read_tokens.map(|value| value as i64)),
cache_write_tokens: usage
.and_then(|usage| usage.cache_write_tokens.map(|value| value as i64)),
reasoning_tokens: usage
.and_then(|usage| usage.reasoning_tokens.map(|value| value as i64)),
},
))
}
pub fn token_delta(tokens: u64) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::TokenDelta(
pb::TokenDeltaUpdate {
tokens: tokens.min(i32::MAX as u64) as i32,
},
))
}
pub fn summary_started() -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::SummaryStarted(
pb::SummaryStartedUpdate {},
))
}
pub fn summary_delta(summary: String) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::Summary(
pb::SummaryUpdate { summary },
))
}
pub fn summary_completed() -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::SummaryCompleted(
pb::SummaryCompletedUpdate { hook_message: None },
))
}
pub fn context_injection_queued(injection_id: String) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ContextInjectionState(
pb::ContextInjectionStateUpdate {
injection_id,
state: Some(pb::ContextInjectionState {
state: Some(pb::context_injection_state::State::Queued(
pb::ContextInjectionQueued {},
)),
}),
},
))
}
pub fn context_injection_rejected(injection_id: String, reason: String) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ContextInjectionState(
pb::ContextInjectionStateUpdate {
injection_id,
state: Some(pb::ContextInjectionState {
state: Some(pb::context_injection_state::State::Rejected(
pb::ContextInjectionRejected { reason },
)),
}),
},
))
}
pub fn context_injection_delivered(
injection_id: String,
delivery_batch_id: String,
delivered_at_ms: i64,
) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ContextInjectionState(
pb::ContextInjectionStateUpdate {
injection_id,
state: Some(pb::ContextInjectionState {
state: Some(pb::context_injection_state::State::Delivered(
pb::ContextInjectionDelivered {
step: 0,
delivery_batch_id,
delivered_at_ms,
},
)),
}),
},
))
}
pub fn user_message_appended(user_message: pb::UserMessage) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::UserMessageAppended(
pb::UserMessageAppendedUpdate {
user_message: Some(user_message),
},
))
}
pub fn server_interaction(message: pb::interaction_update::Message) -> pb::AgentServerMessage {
pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::InteractionUpdate(
pb::InteractionUpdate {
message: Some(message),
},
)),
}
}
+233
View File
@@ -0,0 +1,233 @@
//! Encodes and decodes Cursor JSON streaming payloads.
use crate::{Error, Result};
#[derive(Debug, PartialEq)]
pub(crate) enum StringFieldEvent {
Delta { name: String, text: String },
End { name: String },
}
#[derive(Default)]
pub(crate) struct JsonStringFields {
state: State,
key: String,
string: JsonString,
skipped: SkippedValue,
}
#[derive(Default)]
enum State {
#[default]
Object,
Key,
KeyString,
Colon,
Value,
ValueString,
SkipValue,
AfterValue,
Done,
}
impl JsonStringFields {
pub fn push(&mut self, input: &str) -> Result<Vec<StringFieldEvent>> {
let mut events = Vec::new();
for character in input.chars() {
self.consume(character, &mut events)?;
}
Ok(events)
}
fn consume(&mut self, character: char, events: &mut Vec<StringFieldEvent>) -> Result<()> {
match self.state {
State::Object => match character {
'{' => self.state = State::Key,
value if value.is_whitespace() => {}
_ => return Err(protocol("tool arguments must start with an object")),
},
State::Key => match character {
'"' => {
self.key.clear();
self.string.clear();
self.state = State::KeyString;
}
'}' => self.state = State::Done,
value if value.is_whitespace() => {}
_ => return Err(protocol("expected a tool argument name")),
},
State::KeyString => match self.string.push(character)? {
StringStep::Text(text) => self.key.push_str(&text),
StringStep::End => self.state = State::Colon,
StringStep::Pending => {}
},
State::Colon => match character {
':' => self.state = State::Value,
value if value.is_whitespace() => {}
_ => return Err(protocol("expected ':' after tool argument name")),
},
State::Value => match character {
'"' => {
self.string.clear();
self.state = State::ValueString;
}
value if value.is_whitespace() => {}
value => {
self.skipped.start(value);
self.state = State::SkipValue;
}
},
State::ValueString => match self.string.push(character)? {
StringStep::Text(text) => push_delta(events, &self.key, text),
StringStep::End => {
events.push(StringFieldEvent::End {
name: self.key.clone(),
});
self.state = State::AfterValue;
}
StringStep::Pending => {}
},
State::SkipValue => {
if let Some(terminal) = self.skipped.push(character) {
self.state = match terminal {
',' => State::Key,
'}' => State::Done,
_ => return Err(protocol("invalid skipped JSON value terminator")),
};
}
}
State::AfterValue => match character {
',' => self.state = State::Key,
'}' => self.state = State::Done,
value if value.is_whitespace() => {}
_ => return Err(protocol("expected ',' after tool argument value")),
},
State::Done if character.is_whitespace() => {}
State::Done => return Err(protocol("data after tool arguments object")),
}
Ok(())
}
}
fn push_delta(events: &mut Vec<StringFieldEvent>, name: &str, text: String) {
if let Some(StringFieldEvent::Delta {
name: previous_name,
text: previous_text,
}) = events.last_mut()
{
if previous_name == name {
previous_text.push_str(&text);
return;
}
}
events.push(StringFieldEvent::Delta {
name: name.into(),
text,
});
}
#[derive(Default)]
struct JsonString {
escape: String,
}
enum StringStep {
Text(String),
End,
Pending,
}
impl JsonString {
fn clear(&mut self) {
self.escape.clear();
}
fn push(&mut self, character: char) -> Result<StringStep> {
if self.escape.is_empty() {
return match character {
'"' => Ok(StringStep::End),
'\\' => {
self.escape.push(character);
Ok(StringStep::Pending)
}
value if value < '\u{20}' => Err(protocol("control character in JSON string")),
value => Ok(StringStep::Text(value.to_string())),
};
}
self.escape.push(character);
let complete = match self.escape.as_bytes() {
[b'\\', b'u', a, b, c, d]
if [a, b, c, d].iter().all(|value| value.is_ascii_hexdigit()) =>
{
let code = u16::from_str_radix(&self.escape[2..], 16)
.map_err(|_| protocol("invalid JSON unicode escape"))?;
!(0xD800..=0xDBFF).contains(&code)
}
[b'\\', b'u', ..] if self.escape.len() < 6 => false,
[b'\\', b'u', a, b, c, d, b'\\', b'u', e, f, g, h]
if [a, b, c, d, e, f, g, h]
.iter()
.all(|value| value.is_ascii_hexdigit()) =>
{
true
}
[b'\\', b'u', ..] if self.escape.len() < 12 => false,
[b'\\', b'"' | b'\\' | b'/' | b'b' | b'f' | b'n' | b'r' | b't'] => true,
[b'\\'] => false,
_ => return Err(protocol("invalid JSON string escape")),
};
if !complete {
return Ok(StringStep::Pending);
}
let quoted = format!("\"{}\"", self.escape);
let decoded: String = serde_json::from_str(&quoted)
.map_err(|error| protocol(&format!("invalid JSON string escape: {error}")))?;
self.escape.clear();
Ok(StringStep::Text(decoded))
}
}
#[derive(Default)]
struct SkippedValue {
depth: usize,
string: bool,
escaped: bool,
}
impl SkippedValue {
fn start(&mut self, first: char) {
*self = Self::default();
self.observe(first);
}
fn push(&mut self, character: char) -> Option<char> {
if !self.string && self.depth == 0 && matches!(character, ',' | '}') {
return Some(character);
}
self.observe(character);
None
}
fn observe(&mut self, character: char) {
if self.string {
if self.escaped {
self.escaped = false;
} else if character == '\\' {
self.escaped = true;
} else if character == '"' {
self.string = false;
}
return;
}
match character {
'"' => self.string = true,
'{' | '[' => self.depth += 1,
'}' | ']' => self.depth = self.depth.saturating_sub(1),
_ => {}
}
}
}
fn protocol(message: &str) -> Error {
Error::Protocol(message.into())
}
+6
View File
@@ -0,0 +1,6 @@
//! Exposes Cursor wire protocol primitives outside the Tool protocol.
pub mod connect;
pub mod events;
pub mod json_stream;
pub mod proto;
+72
View File
@@ -0,0 +1,72 @@
//! Includes generated Cursor protobuf types.
pub mod agent {
#[allow(clippy::large_enum_variant)]
pub mod v1 {
include!(concat!(env!("OUT_DIR"), "/agent.v1.rs"));
}
}
pub mod aiserver {
pub mod v1 {
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct BidiRequestId {
#[prost(string, tag = "1")]
pub request_id: String,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct BidiAppendRequest {
#[prost(string, tag = "1")]
pub data: String,
#[prost(message, optional, tag = "2")]
pub request_id: Option<BidiRequestId>,
#[prost(int64, tag = "3")]
pub append_seqno: i64,
#[prost(bytes = "vec", tag = "4")]
pub data_binary: Vec<u8>,
}
#[derive(Clone, Copy, PartialEq, ::prost::Message)]
pub struct BidiAppendResponse {}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct CustomErrorDetails {
#[prost(string, tag = "1")]
pub title: String,
#[prost(string, tag = "2")]
pub detail: String,
#[prost(bool, optional, tag = "3")]
pub allow_command_links_potentially_unsafe_please_only_use_for_handwritten_trusted_markdown:
Option<bool>,
#[prost(bool, optional, tag = "4")]
pub is_retryable: Option<bool>,
#[prost(bool, optional, tag = "5")]
pub show_request_id: Option<bool>,
#[prost(bool, optional, tag = "6")]
pub should_show_immediate_error: Option<bool>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct ErrorDetails {
#[prost(enumeration = "error_details::Error", tag = "1")]
pub error: i32,
#[prost(message, optional, tag = "2")]
pub details: Option<CustomErrorDetails>,
#[prost(bool, optional, tag = "3")]
pub is_expected: Option<bool>,
}
pub mod error_details {
#[derive(
Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, ::prost::Enumeration,
)]
#[repr(i32)]
pub enum Error {
Unspecified = 0,
CustomMessage = 29,
ProviderError = 57,
Internal = 59,
}
}
}
}
+312
View File
@@ -0,0 +1,312 @@
//! Implements Cursor account information services.
use axum::{
body::{Body, Bytes},
extract::Extension,
http::{header, Request, Response},
};
use prost::Message;
use serde_json::{Map, Value};
use crate::{api::cursor::proxy, Result};
const LOCAL_AUTH_ID: &str = "local_ultra";
const LOCAL_EMAIL: &str = "cursor@ai.com";
const LOCAL_ULTRA_PLAN_INCLUDED_CENTS: i32 = 20_000;
#[derive(Clone, PartialEq, Message)]
struct GetEmailResponse {
#[prost(string, tag = "1")]
email: String,
#[prost(int32, tag = "2")]
sign_up_type: i32,
}
#[derive(Clone, PartialEq, Message)]
struct GetMeResponse {
#[prost(string, tag = "1")]
auth_id: String,
#[prost(int32, tag = "2")]
user_id: i32,
#[prost(string, optional, tag = "3")]
email: Option<String>,
#[prost(string, optional, tag = "4")]
first_name: Option<String>,
#[prost(string, optional, tag = "5")]
last_name: Option<String>,
#[prost(string, optional, tag = "8")]
created_at: Option<String>,
#[prost(bool, optional, tag = "9")]
is_enterprise_user: Option<bool>,
#[prost(string, optional, tag = "11")]
email_domain_type: Option<String>,
#[prost(string, optional, tag = "12")]
country: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct GetUserProfileResponse {
#[prost(bool, optional, tag = "4")]
public_visibility_allowed: Option<bool>,
#[prost(string, optional, tag = "5")]
max_visibility: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct GetCurrentPeriodUsageResponse {
#[prost(int64, tag = "1")]
billing_cycle_start: i64,
#[prost(int64, tag = "2")]
billing_cycle_end: i64,
#[prost(message, optional, tag = "3")]
plan_usage: Option<PlanUsage>,
#[prost(message, optional, tag = "4")]
spend_limit_usage: Option<SpendLimitUsage>,
#[prost(int32, optional, tag = "5")]
display_threshold: Option<i32>,
#[prost(bool, tag = "6")]
enabled: bool,
#[prost(string, tag = "7")]
display_message: String,
#[prost(string, optional, tag = "11")]
auto_model_selected_display_message: Option<String>,
#[prost(string, optional, tag = "12")]
named_model_selected_display_message: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct PlanUsage {
#[prost(int32, tag = "1")]
total_spend: i32,
#[prost(int32, tag = "2")]
included_spend: i32,
#[prost(int32, tag = "4")]
remaining: i32,
#[prost(int32, tag = "5")]
limit: i32,
#[prost(bool, optional, tag = "6")]
remaining_bonus: Option<bool>,
#[prost(string, optional, tag = "7")]
bonus_tooltip: Option<String>,
#[prost(int32, optional, tag = "8")]
auto_spend: Option<i32>,
#[prost(int32, optional, tag = "9")]
api_spend: Option<i32>,
#[prost(double, optional, tag = "12")]
auto_percent_used: Option<f64>,
#[prost(double, optional, tag = "13")]
api_percent_used: Option<f64>,
#[prost(double, optional, tag = "14")]
total_percent_used: Option<f64>,
}
#[derive(Clone, PartialEq, Message)]
struct SpendLimitUsage {
#[prost(string, tag = "8")]
limit_type: String,
}
#[derive(Clone, PartialEq, Message)]
struct GetUsageLimitStatusAndActiveGrantsResponse {
#[prost(message, optional, tag = "1")]
usage_limit_policy_status: Option<UsageLimitPolicyStatus>,
}
#[derive(Clone, PartialEq, Message)]
struct UsageLimitPolicyStatus {
#[prost(bool, tag = "1")]
is_in_slow_pool: bool,
#[prost(map = "string, string", tag = "5")]
features: std::collections::HashMap<String, String>,
#[prost(bool, tag = "6")]
can_configure_spend_limit: bool,
#[prost(bool, tag = "8")]
has_pending_request: bool,
#[prost(string, repeated, tag = "9")]
allowed_model_ids: Vec<String>,
#[prost(string, repeated, tag = "10")]
allowed_model_tags: Vec<String>,
}
#[derive(Clone, Copy, PartialEq, Message)]
struct Empty {}
pub async fn get_email(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
proto(GetEmailResponse {
email: LOCAL_EMAIL.into(),
sign_up_type: 3,
})
})
.await
}
pub async fn get_me(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
proto(GetMeResponse {
auth_id: LOCAL_AUTH_ID.into(),
user_id: 1,
email: Some(LOCAL_EMAIL.into()),
first_name: Some("Cursor".into()),
last_name: Some("Local".into()),
created_at: Some(chrono::Utc::now().to_rfc3339()),
is_enterprise_user: Some(false),
email_domain_type: Some("personal".into()),
country: Some("US".into()),
})
})
.await
}
pub async fn get_teams(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || proto(Empty {})).await
}
pub async fn get_user_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
proto(GetUserProfileResponse {
public_visibility_allowed: Some(true),
max_visibility: Some("PUBLIC".into()),
})
})
.await
}
pub async fn current_period_usage() -> Result<Response<Body>> {
let now = chrono::Utc::now();
proto(GetCurrentPeriodUsageResponse {
billing_cycle_start: (now - chrono::Duration::days(30)).timestamp_millis(),
billing_cycle_end: (now + chrono::Duration::days(10 * 365)).timestamp_millis(),
plan_usage: Some(PlanUsage {
total_spend: 0,
included_spend: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
limit: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining_bonus: Some(false),
bonus_tooltip: Some("Ultra local account mock is active.".into()),
auto_spend: Some(0),
api_spend: Some(0),
auto_percent_used: Some(0.0),
api_percent_used: Some(0.0),
total_percent_used: Some(0.0),
}),
spend_limit_usage: Some(SpendLimitUsage {
limit_type: "user".into(),
}),
display_threshold: Some(99_999_999),
enabled: true,
display_message: "Ultra plan active".into(),
auto_model_selected_display_message: Some("Ultra plan active".into()),
named_model_selected_display_message: Some("Ultra plan active".into()),
})
}
pub async fn usage_limit_status() -> Result<Response<Body>> {
proto(GetUsageLimitStatusAndActiveGrantsResponse {
usage_limit_policy_status: Some(UsageLimitPolicyStatus {
is_in_slow_pool: false,
features: Default::default(),
can_configure_spend_limit: true,
has_pending_request: false,
allowed_model_ids: Vec::new(),
allowed_model_tags: Vec::new(),
}),
})
}
pub async fn stripe_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => {
let mut profile = serde_json::from_slice::<Map<String, Value>>(&response.body)?;
ultra(&mut profile);
Ok(response.with_body(Bytes::from(serde_json::to_vec(&profile)?)))
}
Ok(response) => {
tracing::warn!(status = %response.status, "Cursor account upstream rejected profile; using local Ultra identity");
json(ultra_profile())
}
Err(error) => {
tracing::warn!(%error, "Cursor account upstream unavailable; using local Ultra identity");
json(ultra_profile())
}
}
}
async fn forward_or(
upstream: proxy::CursorProxy,
request: Request<Body>,
fallback: impl FnOnce() -> Result<Response<Body>>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => Ok(response.into_response()),
Ok(response) => {
tracing::warn!(status = %response.status, "Cursor identity upstream rejected request; using local identity");
fallback()
}
Err(error) => {
tracing::warn!(%error, "Cursor identity upstream unavailable; using local identity");
fallback()
}
}
}
fn proto(message: impl Message) -> Result<Response<Body>> {
response("application/proto", message.encode_to_vec())
}
fn json(value: Value) -> Result<Response<Body>> {
response("application/json", serde_json::to_vec(&value)?)
}
fn response(content_type: &'static str, body: Vec<u8>) -> Result<Response<Body>> {
let length = body.len();
let mut response = Response::new(Body::from(body));
response.headers_mut().insert(
header::CONTENT_TYPE,
axum::http::HeaderValue::from_static(content_type),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
length
.to_string()
.parse()
.expect("body length is always a valid header value"),
);
Ok(response)
}
fn ultra(profile: &mut Map<String, Value>) {
profile.insert("membershipType".into(), Value::String("ultra".into()));
profile.insert(
"individualMembershipType".into(),
Value::String("ultra".into()),
);
profile.insert("subscriptionStatus".into(), Value::String("active".into()));
}
fn ultra_profile() -> Value {
serde_json::json!({
"membershipType": "ultra",
"individualMembershipType": "ultra",
"subscriptionStatus": "active",
"lastPaymentFailed": false,
"pendingCancellationDate": null,
"daysRemainingOnTrial": 0,
"paymentId": LOCAL_AUTH_ID,
"isTeamMember": false
})
}
+176
View File
@@ -0,0 +1,176 @@
//! Implements Cursor analytics endpoints and event handling.
use axum::{
body::{Body, Bytes},
extract::Extension,
http::{header, HeaderValue, Request, Response, StatusCode},
};
use base64::{engine::general_purpose::STANDARD, Engine};
use bytes::{BufMut, BytesMut};
use prost::Message;
use serde_json::{json, Map, Value};
use sha2::{Digest, Sha256};
use crate::{api::cursor::proxy, Error, Result};
pub const BOOTSTRAP_STATSIG_PATH: &str = "/aiserver.v1.AnalyticsService/BootstrapStatsig";
const AGENT_RETRIES_GATE: &str = "nal_agent_retries";
const LOCAL_RULE: &str = "local_enabled";
#[derive(Clone, PartialEq, Message)]
struct BootstrapStatsigResponse {
#[prost(string, tag = "1")]
config: String,
#[prost(uint64, tag = "2")]
generated_at_ms: u64,
}
pub async fn bootstrap_statsig(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => match patch_upstream(response) {
Ok(response) => Ok(response),
Err(error) => {
tracing::warn!(%error, "Cursor Statsig bootstrap was invalid; using local bootstrap");
local_response()
}
},
Ok(response) => {
tracing::warn!(status = %response.status, "Cursor Statsig bootstrap was rejected; using local bootstrap");
local_response()
}
Err(error) => {
tracing::warn!(%error, "Cursor Statsig bootstrap was unavailable; using local bootstrap");
local_response()
}
}
}
fn patch_upstream(response: proxy::BufferedResponse) -> Result<Response<Body>> {
let (framed, payload) = unary_payload(&response.body)?;
let mut message = BootstrapStatsigResponse::decode(payload)?;
let mut config = serde_json::from_str::<Value>(&message.config)?;
enable_agent_retries(&mut config)?;
message.config = serde_json::to_string(&config)?;
Ok(response.with_body(encode_unary(&message, framed)))
}
fn local_response() -> Result<Response<Body>> {
let generated_at_ms = chrono::Utc::now().timestamp_millis() as u64;
let mut config = json!({
"feature_gates": {},
"dynamic_configs": {},
"layer_configs": {},
"user": {
"userID": "local_ultra",
"customIDs": { "localUserID": "local_ultra" }
},
"has_updates": true,
"hash_used": "none",
"sdkParams": {
"stableID": "local_ultra",
"disableDiagnosticsLogging": true
},
"time": generated_at_ms
});
enable_agent_retries(&mut config)?;
let message = BootstrapStatsigResponse {
config: serde_json::to_string(&config)?,
generated_at_ms,
};
let body = message.encode_to_vec();
let mut response = Response::new(Body::from(body.clone()));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
body.len()
.to_string()
.parse()
.expect("body length is a valid header value"),
);
Ok(response)
}
fn enable_agent_retries(config: &mut Value) -> Result<()> {
let gate_key = statsig_key(config, AGENT_RETRIES_GATE);
let root = config
.as_object_mut()
.ok_or_else(|| Error::Protocol("Statsig bootstrap config must be an object".into()))?;
let gates = root
.entry("feature_gates")
.or_insert_with(|| Value::Object(Map::new()))
.as_object_mut()
.ok_or_else(|| Error::Protocol("Statsig feature_gates must be an object".into()))?;
gates.insert(gate_key.clone(), enabled_gate(&gate_key));
Ok(())
}
fn statsig_key(config: &Value, name: &str) -> String {
match config.get("hash_used").and_then(Value::as_str) {
Some("djb2") => djb2(name),
Some("sha256") => STANDARD.encode(Sha256::digest(name.as_bytes())),
_ => name.to_owned(),
}
}
fn djb2(value: &str) -> String {
value
.encode_utf16()
.fold(0_u32, |hash, character| {
hash.wrapping_mul(31).wrapping_add(u32::from(character))
})
.to_string()
}
fn enabled_gate(name: &str) -> Value {
json!({
"name": name,
"value": true,
"rule_id": LOCAL_RULE,
"ruleID": LOCAL_RULE,
"group_name": LOCAL_RULE,
"groupName": LOCAL_RULE,
"secondary_exposures": [],
"secondaryExposures": [],
"undelegated_secondary_exposures": [],
"undelegatedSecondaryExposures": [],
"is_device_based": false,
"isDeviceBased": false,
"id_type": "userID",
"idType": "userID"
})
}
fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
if body.len() < 5 {
return Ok((false, body));
}
let flags = body[0];
let length = u32::from_be_bytes([body[1], body[2], body[3], body[4]]) as usize;
if length != body.len() - 5 {
return Ok((false, body));
}
if flags != 0 {
return Err(Error::Protocol(format!(
"cannot patch compressed or terminal Statsig frame: flags={flags}"
)));
}
Ok((true, &body[5..]))
}
fn encode_unary(message: &impl Message, framed: bool) -> Bytes {
let payload = message.encode_to_vec();
if !framed {
return Bytes::from(payload);
}
let mut output = BytesMut::with_capacity(5 + payload.len());
output.put_u8(0);
output.put_u32(payload.len() as u32);
output.extend_from_slice(&payload);
output.freeze()
}
+320
View File
@@ -0,0 +1,320 @@
//! Synchronizes content-addressed blobs with Cursor.
use std::{
collections::{HashMap, HashSet},
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
time::Duration,
};
use tokio::sync::{oneshot, Mutex};
use crate::{
cursor::protocol::proto::agent::v1 as pb,
cursor::services::observability::CursorTraceRecorder,
cursor::transport::TransportHandle,
store::{BlobEdge, BlobId, Store},
Error, Result,
};
type BlobSetSender = oneshot::Sender<Result<()>>;
#[derive(Clone)]
pub struct BlobSynchronizer {
inner: Arc<Inner>,
}
struct Inner {
request_id: String,
store: Store,
handle: TransportHandle,
next_id: AtomicU32,
set_requests: Mutex<HashMap<u32, PendingSet>>,
acked_blobs: Mutex<HashSet<BlobId>>,
get_requests: Mutex<HashMap<u32, PendingGet>>,
}
struct PendingSet {
blob_id: BlobId,
sent_at: std::time::Instant,
result: BlobSetSender,
}
struct PendingGet {
blob_id: BlobId,
result: oneshot::Sender<Result<Option<Vec<u8>>>>,
}
impl BlobSynchronizer {
pub fn new(request_id: String, store: Store, handle: TransportHandle) -> Self {
Self {
inner: Arc::new(Inner {
request_id,
store,
handle,
next_id: AtomicU32::new(1),
set_requests: Mutex::new(HashMap::new()),
acked_blobs: Mutex::new(HashSet::new()),
get_requests: Mutex::new(HashMap::new()),
}),
}
}
pub fn request_id(&self) -> &str {
&self.inner.request_id
}
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
self.inner.handle.trace()
}
pub async fn persist(&self, data: &[u8], edges: &[BlobEdge]) -> Result<BlobId> {
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;
}
result?;
Ok(id)
}
async fn ensure_set(&self, blob_id: &BlobId, data: &[u8]) -> Result<()> {
if self.inner.acked_blobs.lock().await.contains(blob_id) {
return Ok(());
}
let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
let (sender, receiver) = oneshot::channel();
self.inner.set_requests.lock().await.insert(
id,
PendingSet {
blob_id: blob_id.clone(),
sent_at: std::time::Instant::now(),
result: sender,
},
);
if let Err(error) = self.inner.handle.emit(&pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::KvServerMessage(
pb::KvServerMessage {
id,
span_context: None,
message: Some(pb::kv_server_message::Message::SetBlobArgs(
pb::SetBlobArgs {
blob_id: blob_id.as_bytes().to_vec(),
blob_data: data.to_vec(),
},
)),
},
)),
}) {
self.inner.set_requests.lock().await.remove(&id);
return Err(error);
}
let cancellation = self.inner.handle.disconnect_token();
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.set_requests.lock().await.remove(&id);
}
result
}
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;
}
return Ok(Some(data));
}
let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
let (sender, receiver) = oneshot::channel();
self.inner.get_requests.lock().await.insert(
id,
PendingGet {
blob_id: blob_id.clone(),
result: sender,
},
);
self.inner.handle.emit(&pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::KvServerMessage(
pb::KvServerMessage {
id,
span_context: None,
message: Some(pb::kv_server_message::Message::GetBlobArgs(
pb::GetBlobArgs {
blob_id: blob_id.as_bytes().to_vec(),
},
)),
},
)),
})?;
let cancellation = self.inner.handle.disconnect_token();
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.get_requests.lock().await.remove(&id);
}
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;
}
Ok(None) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "missing",
}),
)
.await;
}
Err(error) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "error",
"error": error.to_string(),
}),
)
.await;
}
}
}
result
}
pub async fn cache_received(&self, blob_id: &BlobId, data: &[u8]) -> Result<()> {
let actual = BlobId::digest(data);
if actual != *blob_id {
return Err(Error::Protocol(format!(
"received Blob hash mismatch: expected {}, got {}",
blob_id.to_base64(),
actual.to_base64()
)));
}
self.inner.store.put_blob(data, &[]).await?;
Ok(())
}
pub async fn handle_client(&self, message: pb::KvClientMessage) -> Result<()> {
match message.message {
Some(pb::kv_client_message::Message::SetBlobResult(result)) => {
if let Some(pending) = self.inner.set_requests.lock().await.remove(&message.id) {
if let Some(error) = result.error {
tracing::error!(
request_id = self.request_id(),
kv_id = message.id,
blob_id = pending.blob_id.to_base64(),
error = error.message,
"Cursor rejected Blob SET"
);
let _ = pending.result.send(Err(Error::Protocol(format!(
"KV SET {}: {}",
pending.blob_id.to_base64(),
error.message
))));
} else {
tracing::debug!(
request_id = self.request_id(),
kv_id = message.id,
blob_id = pending.blob_id.to_base64(),
elapsed_ms = pending.sent_at.elapsed().as_millis(),
"Cursor acknowledged Blob SET"
);
self.inner.acked_blobs.lock().await.insert(pending.blob_id);
let _ = pending.result.send(Ok(()));
}
} else {
tracing::warn!(
request_id = self.request_id(),
kv_id = message.id,
"unknown Cursor Blob SET acknowledgement"
);
}
}
Some(pb::kv_client_message::Message::GetBlobResult(result)) => {
if let Some(pending) = self.inner.get_requests.lock().await.remove(&message.id) {
let value = if let Some(error) = result.error {
Err(Error::Protocol(format!("KV GET: {}", error.message)))
} else if let Some(data) = result.blob_data {
let actual = BlobId::digest(&data);
if actual != pending.blob_id {
Err(Error::Protocol(format!(
"KV GET Blob hash mismatch: expected {}, got {}",
pending.blob_id.to_base64(),
actual.to_base64()
)))
} else {
self.inner.store.put_blob(&data, &[]).await?;
Ok(Some(data))
}
} else {
Ok(None)
};
let _ = pending.result.send(value);
} else {
tracing::warn!(
request_id = self.request_id(),
kv_id = message.id,
"unknown Cursor Blob GET response"
);
}
}
None => {}
}
Ok(())
}
}
+197
View File
@@ -0,0 +1,197 @@
//! Hydrates request context blobs supplied by Cursor.
use std::{sync::Arc, time::Duration};
use prost::Message;
use tokio::sync::{oneshot, Mutex};
use crate::{
cursor::{protocol::proto::agent::v1 as pb, transport::TransportHandle},
store::{BlobId, Store},
Error, Result,
};
type ContextSender = oneshot::Sender<Result<pb::RequestContext>>;
#[derive(Clone)]
pub(crate) struct RequestContextSynchronizer {
handle: TransportHandle,
store: Store,
pending: Arc<Mutex<Option<ContextSender>>>,
}
impl RequestContextSynchronizer {
pub(crate) fn new(handle: TransportHandle, store: Store) -> Self {
Self {
handle,
store,
pending: Arc::new(Mutex::new(None)),
}
}
pub(crate) async fn refresh_if_missing(
&self,
references: &pb::RequestContextPartReferences,
conversation_id: &str,
) -> Result<Option<pb::RequestContext>> {
if !self.has_missing_part(references).await? {
return Ok(None);
}
let context = self.load(conversation_id).await?;
self.cache_parts(&context).await?;
Ok(Some(context))
}
pub(crate) async fn get(&self, id: &BlobId) -> Result<Option<Vec<u8>>> {
self.store.get_blob(id).await
}
pub(crate) async fn load(&self, conversation_id: &str) -> Result<pb::RequestContext> {
let (sender, receiver) = oneshot::channel();
let mut pending = self.pending.lock().await;
if pending.is_some() {
return Err(Error::Protocol(
"Cursor request context is already being loaded".into(),
));
}
*pending = Some(sender);
drop(pending);
tracing::info!(
request_id = self.handle.request_id(),
conversation_id,
"requesting uncached Cursor context"
);
if let Err(error) = self.handle.emit(&pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ExecServerMessage(
pb::ExecServerMessage {
id: 0,
message: Some(pb::exec_server_message::Message::RequestContextArgs(
pb::RequestContextArgs {
notes_session_id: Some(conversation_id.into()),
..Default::default()
},
)),
..Default::default()
},
)),
}) {
self.pending.lock().await.take();
return Err(error);
}
let cancellation = self.handle.disconnect_token();
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("request context response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol("request context timed out".into())),
};
if result.is_err() {
self.pending.lock().await.take();
}
result
}
pub(crate) async fn handle_client(&self, message: &pb::ExecClientMessage) -> bool {
if message.id != 0 {
return false;
}
let Some(pb::exec_client_message::Message::RequestContextResult(result)) =
message.message.as_ref()
else {
return false;
};
let Some(sender) = self.pending.lock().await.take() else {
tracing::warn!(
request_id = self.handle.request_id(),
"unexpected Cursor request context result"
);
return true;
};
use pb::request_context_result::Result as ContextResult;
let result = match result.result.as_ref() {
Some(ContextResult::Success(success)) => success
.request_context
.clone()
.ok_or_else(|| Error::Protocol("Cursor returned empty request context".into())),
Some(ContextResult::Error(error)) => Err(Error::Protocol(format!(
"Cursor request context failed: {}",
error.error
))),
Some(ContextResult::Rejected(rejected)) => Err(Error::Protocol(format!(
"Cursor rejected request context: {}",
rejected.reason
))),
None => Err(Error::Protocol(
"Cursor returned no request context result".into(),
)),
};
let _ = sender.send(result);
true
}
pub(crate) async fn handle_stream_close(&self, id: u32) -> bool {
id == 0 && self.pending.lock().await.is_some()
}
pub(crate) async fn handle_throw(&self, id: u32, message: String) -> bool {
let sender = if id == 0 {
self.pending.lock().await.take()
} else {
None
};
let Some(sender) = sender else { return false };
let _ = sender.send(Err(Error::Protocol(message)));
true
}
async fn has_missing_part(&self, parts: &pb::RequestContextPartReferences) -> Result<bool> {
for raw_id in [
parts.rules_blob_id.as_slice(),
parts.skills_blob_id.as_slice(),
parts.subagents_blob_id.as_slice(),
parts.mcps_blob_id.as_slice(),
] {
if raw_id.is_empty() {
continue;
}
let id = BlobId::from_bytes(raw_id)?;
if self.store.get_blob(&id).await?.is_none() {
return Ok(true);
}
}
Ok(false)
}
async fn cache_parts(&self, context: &pb::RequestContext) -> Result<()> {
self.cache_part(&pb::RequestContextRulesPart {
rules: context.rules.clone(),
non_file_rules: context.non_file_rules.clone(),
cloud_rule: context.cloud_rule.clone(),
})
.await?;
self.cache_part(&pb::RequestContextSkillsPart {
agent_skills: context.agent_skills.clone(),
skill_options: context.skill_options.clone(),
})
.await?;
self.cache_part(&pb::RequestContextSubagentsPart {
custom_subagents: context.custom_subagents.clone(),
})
.await?;
self.cache_part(&pb::RequestContextMcpsPart {
tools: context.tools.clone(),
mcp_instructions: context.mcp_instructions.clone(),
mcp_file_system_options: context.mcp_file_system_options.clone(),
mcp_meta_tool_options: context.mcp_meta_tool_options.clone(),
})
.await
}
async fn cache_part<T: Message>(&self, part: &T) -> Result<()> {
let data = part.encode_to_vec();
self.store.put_blob(&data, &[]).await?;
Ok(())
}
}
+10
View File
@@ -0,0 +1,10 @@
//! Exposes Cursor services outside the Agent loop.
pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod context_sync;
pub mod model_catalog;
pub mod observability;
pub mod tab;
pub mod usage;
+520
View File
@@ -0,0 +1,520 @@
//! Publishes the configured model catalog to Cursor.
use axum::{
body::{Body, Bytes},
extract::{Extension, State},
http::{header, HeaderValue, Request, Response, StatusCode},
};
use bytes::{BufMut, BytesMut};
use prost::Message;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
Error, Result,
};
#[derive(Clone, PartialEq, Message)]
struct AvailableModelsAddition {
#[prost(string, repeated, tag = "1")]
model_names: Vec<String>,
#[prost(message, repeated, tag = "2")]
models: Vec<AvailableModel>,
}
#[derive(Clone, PartialEq, Message)]
struct AvailableModel {
#[prost(string, tag = "1")]
name: String,
#[prost(bool, tag = "2")]
default_on: bool,
#[prost(bool, optional, tag = "5")]
supports_agent: Option<bool>,
#[prost(int32, optional, tag = "6")]
degradation_status: Option<i32>,
#[prost(message, optional, tag = "8")]
tooltip_data: Option<TooltipData>,
#[prost(bool, optional, tag = "9")]
supports_thinking: Option<bool>,
#[prost(bool, optional, tag = "10")]
supports_images: Option<bool>,
#[prost(bool, optional, tag = "14")]
supports_max_mode: Option<bool>,
#[prost(string, optional, tag = "17")]
client_display_name: Option<String>,
#[prost(string, optional, tag = "18")]
server_model_name: Option<String>,
#[prost(bool, optional, tag = "19")]
supports_non_max_mode: Option<bool>,
#[prost(message, optional, tag = "20")]
tooltip_data_for_max_mode: Option<TooltipData>,
#[prost(bool, optional, tag = "21")]
is_recommended_for_background_composer: Option<bool>,
#[prost(bool, optional, tag = "22")]
supports_plan_mode: Option<bool>,
#[prost(string, optional, tag = "24")]
inputbox_short_model_name: Option<String>,
#[prost(bool, optional, tag = "25")]
supports_sandboxing: Option<bool>,
#[prost(bool, optional, tag = "26")]
supports_cmd_k: Option<bool>,
#[prost(message, repeated, tag = "29")]
parameter_definitions: Vec<ModelParameterDefinition>,
#[prost(message, repeated, tag = "30")]
variants: Vec<ModelVariant>,
#[prost(string, repeated, tag = "36")]
legacy_slugs: Vec<String>,
#[prost(int32, optional, tag = "38")]
named_model_section_index: Option<i32>,
#[prost(string, optional, tag = "41")]
vendor_name: Option<String>,
#[prost(message, optional, tag = "42")]
vendor: Option<AvailableModelVendor>,
#[prost(message, repeated, tag = "48")]
model_picker_badges: Vec<ModelPickerBadge>,
}
#[derive(Clone, PartialEq, Message)]
struct TooltipData {
#[prost(string, optional, tag = "7")]
markdown_content: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct ModelParameterDefinition {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
name: String,
#[prost(string, optional, tag = "3")]
markdown_tooltip: Option<String>,
#[prost(message, optional, tag = "4")]
parameter_type: Option<ModelParameterType>,
#[prost(bool, optional, tag = "5")]
is_cycleable_by_hotkey: Option<bool>,
}
#[derive(Clone, PartialEq, Message)]
struct ModelParameterType {
#[prost(message, optional, tag = "1")]
boolean_parameter: Option<BooleanParameter>,
#[prost(message, optional, tag = "2")]
enum_parameter: Option<EnumParameter>,
}
#[derive(Clone, PartialEq, Message)]
struct BooleanParameter {
#[prost(message, repeated, tag = "1")]
values: Vec<BooleanParameterValue>,
}
#[derive(Clone, PartialEq, Message)]
struct BooleanParameterValue {
#[prost(string, tag = "1")]
value: String,
#[prost(string, optional, tag = "2")]
display_name: Option<String>,
#[prost(bool, optional, tag = "3")]
increases_model_cost: Option<bool>,
}
#[derive(Clone, PartialEq, Message)]
struct EnumParameter {
#[prost(message, repeated, tag = "1")]
values: Vec<EnumParameterValue>,
}
#[derive(Clone, PartialEq, Message)]
struct EnumParameterValue {
#[prost(string, tag = "1")]
value: String,
#[prost(string, optional, tag = "2")]
display_name: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct ModelVariant {
#[prost(message, repeated, tag = "1")]
parameter_values: Vec<ModelParameterValue>,
#[prost(string, tag = "2")]
display_name: String,
#[prost(bool, tag = "3")]
is_max_mode: bool,
#[prost(bool, optional, tag = "4")]
is_default_max_config: Option<bool>,
#[prost(bool, optional, tag = "5")]
is_default_non_max_config: Option<bool>,
#[prost(message, optional, tag = "6")]
tooltip_data: Option<TooltipData>,
#[prost(string, optional, tag = "8")]
display_name_outside_picker: Option<String>,
#[prost(string, optional, tag = "9")]
variant_string_representation: Option<String>,
#[prost(string, optional, tag = "11")]
legacy_slug: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct ModelParameterValue {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
value: String,
}
#[derive(Clone, PartialEq, Message)]
struct ModelPickerBadge {
#[prost(string, tag = "1")]
label: String,
#[prost(int32, tag = "2")]
variant: i32,
#[prost(bool, tag = "3")]
dismiss_on_selection: bool,
}
#[derive(Clone, PartialEq, Message)]
struct AvailableModelVendor {
#[prost(int32, tag = "1")]
id: i32,
#[prost(string, tag = "2")]
display_name: String,
}
#[derive(Clone, PartialEq, Message)]
struct UsableModelsAddition {
#[prost(message, repeated, tag = "1")]
models: Vec<agent::ModelDetails>,
}
const CONTEXTS: [(&str, &str); 4] = [
("200k", "200K"),
("356k", "356K"),
("800k", "800K"),
("1m", "1M"),
];
const EFFORTS: [(&str, &str); 5] = [
("low", "Low"),
("medium", "Medium"),
("high", "High"),
("xhigh", "Extra High"),
("max", "Max"),
];
const DEFAULT_CONTEXT: &str = "200k";
fn context_options(model: &ModelConfig) -> Vec<(String, String)> {
let mut contexts = CONTEXTS
.into_iter()
.map(|(value, display_name)| (value.to_owned(), display_name.to_owned()))
.collect::<Vec<_>>();
if let Some(tokens) = model.context_window_tokens {
let value = tokens.to_string();
let duplicate = contexts
.iter()
.any(|(existing, _)| parse_token_count(existing) == Some(tokens));
if !duplicate {
contexts.push((value, format!("{} (Custom)", format_token_count(tokens))));
}
}
contexts
}
pub async fn available_models(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
tracing::info!(
model_count = models.len(),
"appending BYOK models to Cursor AvailableModels"
);
let available_models = models.iter().map(available_model).collect::<Vec<_>>();
let local = AvailableModelsAddition {
model_names: models
.iter()
.map(|model| model.model_hash.clone())
.collect(),
models: available_models,
}
.encode_to_vec();
match proxy::forward_buffered(&proxy, request).await {
Ok(upstream) => merge_response(upstream, local),
Err(error) => {
tracing::warn!(%error, "Cursor AvailableModels upstream unavailable; using local catalog");
Ok(local_response(local))
}
}
}
pub async fn usable_models(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
tracing::info!(
model_count = models.len(),
"appending BYOK models to Cursor GetUsableModels"
);
let local = UsableModelsAddition {
models: models.iter().map(usable_model).collect(),
}
.encode_to_vec();
match proxy::forward_buffered(&proxy, request).await {
Ok(upstream) => merge_response(upstream, local),
Err(error) => {
tracing::warn!(%error, "Cursor GetUsableModels upstream unavailable; using local catalog");
Ok(local_response(local))
}
}
}
fn merge_response(upstream: proxy::BufferedResponse, extra: Vec<u8>) -> Result<Response<Body>> {
if !upstream.status.is_success() {
tracing::warn!(status = %upstream.status, "Cursor model catalog upstream rejected request; using local catalog");
return Ok(local_response(extra));
}
let (framed, payload) = unary_payload(&upstream.body)?;
let body = if framed {
let mut merged = BytesMut::with_capacity(5 + payload.len() + extra.len());
merged.put_u8(0);
merged.put_u32((payload.len() + extra.len()) as u32);
merged.extend_from_slice(payload);
merged.extend_from_slice(&extra);
merged.freeze()
} else {
let mut merged = BytesMut::with_capacity(payload.len() + extra.len());
merged.extend_from_slice(payload);
merged.extend_from_slice(&extra);
merged.freeze()
};
Ok(upstream.with_body(body))
}
fn local_response(body: Vec<u8>) -> Response<Body> {
let mut response = Response::new(Body::from(body));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
response
}
fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
if body.len() < 5 {
return Ok((false, body));
}
let flags = body[0];
let length = u32::from_be_bytes([body[1], body[2], body[3], body[4]]) as usize;
if length != body.len() - 5 {
return Ok((false, body));
}
if flags != 0 {
return Err(Error::Protocol(format!(
"cannot merge compressed or terminal model catalog frame: flags={flags}"
)));
}
Ok((true, &body[5..]))
}
fn available_model(model: &ModelConfig) -> AvailableModel {
let contexts = context_options(model);
let variants = model_variants(model, &contexts);
let legacy_slugs = variants
.iter()
.filter_map(|variant| variant.legacy_slug.clone())
.collect();
let tooltip = model_tooltip(model);
AvailableModel {
name: model.model_hash.clone(),
default_on: true,
supports_agent: Some(true),
degradation_status: Some(0),
tooltip_data: Some(tooltip.clone()),
supports_thinking: Some(true),
supports_images: Some(true),
supports_max_mode: Some(true),
client_display_name: Some(model.display_name.clone()),
server_model_name: Some(model.model_hash.clone()),
supports_non_max_mode: Some(true),
tooltip_data_for_max_mode: Some(tooltip),
is_recommended_for_background_composer: Some(false),
supports_plan_mode: Some(true),
inputbox_short_model_name: Some(model.display_name.clone()),
supports_sandboxing: Some(true),
supports_cmd_k: Some(false),
parameter_definitions: model_parameters(&contexts),
variants,
legacy_slugs,
named_model_section_index: Some(1),
vendor_name: Some("cursor".into()),
vendor: Some(AvailableModelVendor {
id: 6,
display_name: "Cursor".into(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: match model.model_type {
ModelType::OpenAi => "OpenAI".into(),
ModelType::Anthropic => "Anthropic".into(),
},
variant: 1,
dismiss_on_selection: false,
}],
}
}
fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefinition> {
vec![
ModelParameterDefinition {
id: "context".into(),
name: "Context".into(),
markdown_tooltip: Some("Context size used to trigger conversation compaction.".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: None,
enum_parameter: Some(EnumParameter {
values: contexts
.iter()
.map(|(value, display_name)| EnumParameterValue {
value: value.clone(),
display_name: Some(display_name.clone()),
})
.collect(),
}),
}),
is_cycleable_by_hotkey: Some(false),
},
ModelParameterDefinition {
id: "reasoning".into(),
name: "Effort".into(),
markdown_tooltip: Some("Effort the model uses to generate its response.".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: None,
enum_parameter: Some(EnumParameter {
values: EFFORTS
.into_iter()
.map(|(value, display_name)| EnumParameterValue {
value: value.into(),
display_name: Some(display_name.into()),
})
.collect(),
}),
}),
is_cycleable_by_hotkey: Some(true),
},
ModelParameterDefinition {
id: "fast".into(),
name: "Fast".into(),
markdown_tooltip: Some("Significantly faster but consumes more usage".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: Some(BooleanParameter {
values: vec![
BooleanParameterValue {
value: "false".into(),
display_name: None,
increases_model_cost: None,
},
BooleanParameterValue {
value: "true".into(),
display_name: Some("Fast".into()),
increases_model_cost: Some(true),
},
],
}),
enum_parameter: None,
}),
is_cycleable_by_hotkey: Some(false),
},
]
}
fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<ModelVariant> {
let mut variants = Vec::with_capacity(contexts.len() * EFFORTS.len() * 2);
for (context, context_name) in contexts {
for (effort, effort_name) in EFFORTS {
for fast in [false, true] {
variants.push(model_variant(
model,
context,
context_name,
effort,
effort_name,
fast,
));
}
}
}
variants
}
fn model_variant(
model: &ModelConfig,
context: &str,
context_name: &str,
effort: &str,
effort_name: &str,
fast: bool,
) -> ModelVariant {
let mut suffix = Vec::with_capacity(3);
if context != DEFAULT_CONTEXT {
suffix.push(context_name);
}
suffix.push(effort_name);
if fast {
suffix.push("Fast");
}
let suffix = suffix.join(" ");
let display_name = format!(
"{} <span style=\"color: var(--cursor-text-tertiary);\">{suffix}</span>",
model.display_name
);
let is_default = context == DEFAULT_CONTEXT && effort == "high" && !fast;
ModelVariant {
parameter_values: vec![
ModelParameterValue {
id: "context".into(),
value: context.into(),
},
ModelParameterValue {
id: "reasoning".into(),
value: effort.into(),
},
ModelParameterValue {
id: "fast".into(),
value: fast.to_string(),
},
],
display_name: display_name.clone(),
is_max_mode: false,
is_default_max_config: is_default.then_some(true),
is_default_non_max_config: is_default.then_some(true),
tooltip_data: Some(model_tooltip(model)),
display_name_outside_picker: Some(display_name),
variant_string_representation: Some(format!(
"{}[context={context},reasoning={effort},fast={fast}]",
model.model_hash
)),
legacy_slug: Some(format!(
"{}-{context}-{effort}{}",
model.model_hash,
if fast { "-fast" } else { "" }
)),
}
}
fn model_tooltip(model: &ModelConfig) -> TooltipData {
TooltipData {
markdown_content: Some(model.tooltip_data.clone()),
}
}
fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
agent::ModelDetails {
model_id: model.model_hash.clone(),
display_model_id: model.model_hash.clone(),
display_name: model.display_name.clone(),
display_name_short: model.display_name.clone(),
thinking_details: Some(agent::ThinkingDetails::default()),
..Default::default()
}
}
+225
View File
@@ -0,0 +1,225 @@
//! 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(())
}
}
+52
View File
@@ -0,0 +1,52 @@
//! Implements Cursor tab metadata services.
use axum::{
body::Body,
extract::{Extension, State},
http::{Request, Response},
routing::post,
Router,
};
use crate::{api::cursor::proxy, cursor::transport::TransportRegistry, Result};
pub const TAB_PATHS: [&str; 17] = [
"/aiserver.v1.AiService/StreamCpp",
"/aiserver.v1.AiService/StreamNextCursorPrediction",
"/aiserver.v1.AiService/GetCppEditClassification",
"/aiserver.v1.AiService/RefreshTabContext",
"/aiserver.v1.AiService/CppConfig",
"/aiserver.v1.AiService/CppEditHistoryStatus",
"/aiserver.v1.AiService/CppAppend",
"/aiserver.v1.AiService/CppEditHistoryAppend",
"/aiserver.v1.AiService/ReportAiCodeChangeMetrics",
"/aiserver.v1.AiService/WriteGitCommitMessage",
"/aiserver.v1.AiService/WriteGitBranchName",
"/aiserver.v1.CppService/AvailableModels",
"/aiserver.v1.CppService/RecordCppFate",
"/aiserver.v1.FileSyncService/FSSyncFile",
"/aiserver.v1.FileSyncService/FSIsEnabledForUser",
"/aiserver.v1.FileSyncService/FSConfig",
"/aiserver.v1.FileSyncService/FSUploadFile",
];
pub fn is_tab_path(path: &str) -> bool {
TAB_PATHS.contains(&path)
}
pub fn router() -> Router<TransportRegistry> {
TAB_PATHS.into_iter().fold(Router::new(), |router, path| {
router.route(path, post(forward))
})
}
async fn forward(
State(registry): State<TransportRegistry>,
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let settings = registry.store().tab_settings().await?;
match settings.service_url() {
Some(service_url) => proxy::forward_to_service(&upstream, request, service_url).await,
None => proxy::forward(Extension(upstream), request).await,
}
}
+219
View File
@@ -0,0 +1,219 @@
//! Builds Cursor usage and context breakdown data.
use std::collections::HashSet;
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{CanonicalMessage, ContentPart, MessageContent, Origin, ToolDefinition},
Result,
};
const CATEGORIES: [(&str, &str); 8] = [
("system_prompt", "System prompt"),
("tools", "Tool definitions"),
("rules", "Rules"),
("skills", "Skills"),
("mcp", "MCP & dynamic tools"),
("subagents", "Subagent definitions"),
("summarized_conversation", "Summarized conversation"),
("conversation", "Conversation"),
];
const EASTER_EGG_CATEGORY: (&str, &str) = ("leookun", "@leookun stole 1 token 😂");
const SYSTEM: usize = 0;
const TOOLS: usize = 1;
const RULES: usize = 2;
const SKILLS: usize = 3;
const MCP: usize = 4;
const SUBAGENTS: usize = 5;
const SUMMARY: usize = 6;
const CONVERSATION: usize = 7;
#[derive(Clone, Copy, Default)]
struct Measure {
characters: u64,
token_units: u64,
}
impl Measure {
fn add(&mut self, text: &str) {
self.characters += text.encode_utf16().count() as u64;
let mut units = 0_u64;
for character in text.chars() {
let width = character.len_utf16() as u64;
units += if character.is_ascii() {
width * 273
} else {
width * 550
};
}
self.token_units += units;
}
fn estimated_tokens(self) -> u64 {
self.token_units.div_ceil(1_000)
}
}
pub(crate) fn breakdown(
used_tokens: u32,
max_tokens: u32,
baseline: Option<&pb::PromptTokenBreakdownSnapshot>,
instructions: &str,
tools: &[ToolDefinition],
dynamic_tools: &HashSet<String>,
messages: &[CanonicalMessage],
) -> Result<pb::PromptTokenBreakdownSnapshot> {
let mut measures = [Measure::default(); 8];
measures[SYSTEM].add(instructions);
for tool in tools {
let encoded = serde_json::to_string(tool)?;
if dynamic_tools.contains(&tool.name) {
measures[MCP].add(&encoded);
} else {
measures[TOOLS].add(&encoded);
}
}
for message in messages {
measure_message(message, &mut measures)?;
}
let mut estimates = [0_u64; 8];
for index in 0..CONVERSATION {
estimates[index] = measures[index].estimated_tokens();
}
if measures[SUMMARY].characters != 0 {
estimates[SUMMARY] = measures[SUMMARY].estimated_tokens();
} else if let Some(summary) = baseline.and_then(|snapshot| {
snapshot
.categories
.iter()
.find(|category| category.id == CATEGORIES[SUMMARY].0)
}) {
measures[SUMMARY].characters = summary.character_count.unwrap_or(0) as u64;
estimates[SUMMARY] = summary.estimated_tokens as u64;
}
let easter_egg_tokens = 1_u64;
let categorized_tokens = used_tokens as u64;
fit_special_estimates(&mut estimates, categorized_tokens);
estimates[CONVERSATION] =
categorized_tokens.saturating_sub(estimates[..CONVERSATION].iter().sum::<u64>());
let mut categories = CATEGORIES
.iter()
.enumerate()
.map(|(index, (id, label))| pb::PromptTokenBreakdownCategory {
id: (*id).into(),
label: (*label).into(),
estimated_tokens: estimates[index].min(u32::MAX as u64) as u32,
character_count: (measures[index].characters != 0)
.then_some(measures[index].characters.min(u32::MAX as u64) as u32),
})
.collect::<Vec<_>>();
categories.push(pb::PromptTokenBreakdownCategory {
id: EASTER_EGG_CATEGORY.0.into(),
label: EASTER_EGG_CATEGORY.1.into(),
estimated_tokens: easter_egg_tokens as u32,
character_count: None,
});
Ok(pb::PromptTokenBreakdownSnapshot {
total_used_tokens: used_tokens,
max_tokens,
categories,
})
}
fn measure_message(message: &CanonicalMessage, measures: &mut [Measure; 8]) -> Result<()> {
match &message.content {
MessageContent::Parts { parts } => {
for part in parts {
if let ContentPart::Text { text } = part {
if message.origin == Origin::Runtime {
measure_runtime(text, measures);
} else {
measures[CONVERSATION].add(text);
}
}
}
}
MessageContent::Assistant {
text,
thinking,
tool_calls,
..
} => {
measures[CONVERSATION].add(text);
measures[CONVERSATION].add(thinking);
measures[CONVERSATION].add(&serde_json::to_string(tool_calls)?);
}
MessageContent::ToolResult(result) => {
measures[CONVERSATION].add(&serde_json::to_string(result)?);
}
}
Ok(())
}
fn measure_runtime(text: &str, measures: &mut [Measure; 8]) {
let mut ranges = Vec::new();
collect_ranges(text, "rules", RULES, &mut ranges);
collect_ranges(text, "rule", RULES, &mut ranges);
collect_ranges(text, "agent_skills", SKILLS, &mut ranges);
collect_ranges(text, "skill", SKILLS, &mut ranges);
collect_ranges(text, "subagents", SUBAGENTS, &mut ranges);
collect_ranges(text, "mcp_meta_tools", MCP, &mut ranges);
collect_ranges(text, "conversation_summary", SUMMARY, &mut ranges);
ranges.sort_by_key(|range| range.0);
let mut cursor = 0;
for (start, end, category) in ranges {
if start < cursor {
continue;
}
measures[CONVERSATION].add(&text[cursor..start]);
measures[category].add(&text[start..end]);
cursor = end;
}
measures[CONVERSATION].add(&text[cursor..]);
}
fn collect_ranges(text: &str, tag: &str, category: usize, output: &mut Vec<(usize, usize, usize)>) {
let opening = format!("<{tag}");
let closing = format!("</{tag}>");
let mut cursor = 0;
while let Some(relative_start) = text[cursor..].find(&opening) {
let start = cursor + relative_start;
let Some(open_end) = text[start..].find('>').map(|offset| start + offset + 1) else {
break;
};
let Some(relative_end) = text[open_end..].find(&closing) else {
break;
};
let end = open_end + relative_end + closing.len();
output.push((start, end, category));
cursor = end;
}
}
fn fit_special_estimates(estimates: &mut [u64; 8], total: u64) {
let special_total = estimates[..CONVERSATION].iter().sum::<u64>();
if special_total <= total || special_total == 0 {
return;
}
let original = *estimates;
let mut assigned = 0;
for index in 0..CONVERSATION {
estimates[index] = original[index].saturating_mul(total) / special_total;
assigned += estimates[index];
}
let mut remainder = total - assigned;
let mut order = (0..CONVERSATION).collect::<Vec<_>>();
order.sort_by_key(|index| {
std::cmp::Reverse(original[*index].saturating_mul(total) % special_total)
});
for index in order {
if remainder == 0 {
break;
}
estimates[index] += 1;
remainder -= 1;
}
}
+93
View File
@@ -0,0 +1,93 @@
//! Encodes and decodes Cursor Tool wire messages.
mod query;
mod render;
mod request;
mod response;
use crate::{
cursor::protocol::{events::server_interaction, proto::agent::v1 as pb},
model::ToolCall,
Error, Result,
};
pub use query::tool_query;
pub(crate) use render::{create_plan_partial, edit_content_delta, edit_path_partial};
pub use render::{dynamic_mcp_placeholder, render_dynamic_mcp, tool_completed};
pub use request::{abort, mcp_request, mcp_state_request, request};
pub(crate) use request::{edit_read_request, json_object_to_prost, mcp_meta_request};
pub use response::{client_event, stream_closed, ClientExecEvent};
use render::{
render_tool_call as render_builtin_tool_call, tool_placeholder as builtin_tool_placeholder,
tool_started as builtin_tool_started,
};
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
match builtin_tool_placeholder(name, call_id) {
Ok(tool) => Ok(tool),
Err(error) if is_unsupported_tool(&error, name) => {
Ok(super::compat::placeholder(name, call_id))
}
Err(error) => Err(error),
}
}
pub fn render_tool_call(call: &ToolCall, completed: bool) -> Result<pb::ToolCall> {
match render_builtin_tool_call(call, completed) {
Ok(tool) => Ok(tool),
Err(error) if is_unsupported_tool(&error, &call.name) => {
Ok(super::compat::render(call, completed))
}
Err(error) => Err(error),
}
}
pub fn tool_started(
call: &ToolCall,
dynamic_mcp: Option<&pb::McpToolDefinition>,
) -> Result<pb::AgentServerMessage> {
match builtin_tool_started(call, dynamic_mcp) {
Ok(message) => Ok(message),
Err(error) if dynamic_mcp.is_none() && is_unsupported_tool(&error, &call.name) => {
Ok(server_interaction(
pb::interaction_update::Message::ToolCallStarted(pb::ToolCallStartedUpdate {
call_id: call.call_id.clone(),
tool_call: Some(super::compat::render(call, false)),
model_call_id: call.model_call_id.clone(),
}),
))
}
Err(error) => Err(error),
}
}
fn is_unsupported_tool(error: &Error, name: &str) -> bool {
matches!(error, Error::Protocol(message) if message == &format!("unsupported tool: {name}"))
}
pub fn arguments_delta(call: &ToolCall, delta: &str) -> Result<pb::AgentServerMessage> {
Ok(server_interaction(
pb::interaction_update::Message::PartialToolCall(pb::PartialToolCallUpdate {
call_id: call.call_id.clone(),
tool_call: Some(tool_placeholder(&call.name, &call.call_id)?),
args_text_delta: delta.into(),
model_call_id: call.model_call_id.clone(),
}),
))
}
pub fn dynamic_mcp_arguments_delta(
call: &ToolCall,
delta: &str,
definition: &pb::McpToolDefinition,
) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::PartialToolCall(
pb::PartialToolCallUpdate {
call_id: call.call_id.clone(),
tool_call: Some(dynamic_mcp_placeholder(definition, &call.call_id)),
args_text_delta: delta.into(),
model_call_id: call.model_call_id.clone(),
},
))
}
+216
View File
@@ -0,0 +1,216 @@
//! Encodes Tool calls as Cursor InteractionQuery messages.
use serde_json::Value;
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolCall, Error, Result};
pub fn tool_query(id: u32, call: &ToolCall) -> Result<pb::AgentServerMessage> {
use pb::interaction_query::Query;
let string = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
.ok_or_else(|| Error::Protocol(format!("{} is missing {name}", call.name)))
};
let optional_string = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
};
let query = match normalized(&call.name).as_str() {
"askquestion" => {
let questions = call
.arguments
.get("questions")
.and_then(Value::as_array)
.into_iter()
.flatten()
.map(|question| -> Result<_> {
let required = |name: &str| {
question
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
.ok_or_else(|| Error::Protocol(format!("question is missing {name}")))
};
let options = question
.get("options")
.and_then(Value::as_array)
.into_iter()
.flatten()
.map(|option| -> Result<_> {
let value = |name: &str| {
option
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
.ok_or_else(|| {
Error::Protocol(format!(
"question option is missing {name}"
))
})
};
Ok(pb::ask_question_args::Option {
id: value("id")?,
label: value("label")?,
})
})
.collect::<Result<Vec<_>>>()?;
Ok(pb::ask_question_args::Question {
id: required("id")?,
prompt: required("prompt")?,
options,
allow_multiple: question
.get("allow_multiple")
.and_then(Value::as_bool)
.unwrap_or(false),
})
})
.collect::<Result<Vec<_>>>()?;
Query::AskQuestionInteractionQuery(pb::AskQuestionInteractionQuery {
args: Some(pb::AskQuestionArgs {
title: optional_string("title").unwrap_or_default(),
questions,
run_async: false,
async_original_tool_call_id: String::new(),
}),
tool_call_id: call.call_id.clone(),
})
}
"websearch" => Query::WebSearchRequestQuery(pb::WebSearchRequestQuery {
args: Some(pb::WebSearchArgs {
search_term: string("search_term")?,
tool_call_id: call.call_id.clone(),
}),
}),
"webfetch" => Query::WebFetchRequestQuery(pb::WebFetchRequestQuery {
args: Some(pb::WebFetchArgs {
url: string("url")?,
tool_call_id: call.call_id.clone(),
}),
skip_approval: false,
smart_mode_approval: smart_mode_approval(
call,
"requestSmartModeApproval",
"smartModeBlockReason",
)?,
}),
"switchmode" => Query::SwitchModeRequestQuery(pb::SwitchModeRequestQuery {
args: Some(pb::SwitchModeArgs {
target_mode_id: string("target_mode_id")?,
explanation: optional_string("explanation"),
tool_call_id: call.call_id.clone(),
}),
}),
"createplan" => {
let todos = call
.arguments
.get("todos")
.and_then(Value::as_array)
.into_iter()
.flatten()
.map(|todo| pb::TodoItem {
id: todo
.get("id")
.and_then(Value::as_str)
.unwrap_or_default()
.into(),
content: todo
.get("content")
.and_then(Value::as_str)
.unwrap_or_default()
.into(),
status: pb::TodoStatus::Pending as i32,
created_at: 0,
updated_at: 0,
dependencies: Vec::new(),
})
.collect();
Query::CreatePlanRequestQuery(pb::CreatePlanRequestQuery {
args: Some(pb::CreatePlanArgs {
plan: string("plan")?,
todos,
overview: string("overview")?,
name: optional_string("name").unwrap_or_default(),
is_project: false,
phases: Vec::new(),
}),
tool_call_id: call.call_id.clone(),
})
}
"generateimage" => Query::GenerateImageRequestQuery(pb::GenerateImageRequestQuery {
args: Some(pb::GenerateImageArgs {
description: string("description")?,
file_path: optional_string("filename"),
reference_image_paths: call
.arguments
.get("reference_image_paths")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_string)
.collect(),
aspect_ratio: optional_string("aspect_ratio"),
}),
tool_call_id: call.call_id.clone(),
}),
"callmcptool"
if optional_string("toolName").is_some_and(|tool| normalized(&tool) == "mcpauth") =>
{
Query::McpAuthRequestQuery(pb::McpAuthRequestQuery {
args: Some(pb::McpAuthArgs {
server_identifier: string("server")?,
tool_call_id: call.call_id.clone(),
}),
})
}
other => {
return Err(Error::Protocol(format!(
"tool {other} is not an InteractionQuery"
)))
}
};
Ok(pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::InteractionQuery(
pb::InteractionQuery {
id,
query: Some(query),
},
)),
})
}
fn smart_mode_approval(
call: &ToolCall,
request_field: &str,
reason_field: &str,
) -> Result<Option<pb::SmartModeApproval>> {
if !call
.arguments
.get(request_field)
.and_then(Value::as_bool)
.unwrap_or(false)
{
return Ok(None);
}
let reason = call
.arguments
.get(reason_field)
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol(format!("{} requires {reason_field}", call.name)))?;
Ok(Some(pb::SmartModeApproval {
request_id: call.call_id.clone(),
reason: reason.to_string(),
}))
}
fn normalized(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
+529
View File
@@ -0,0 +1,529 @@
//! Renders Tool calls and results as Cursor Tool cards.
use serde_json::Value;
use crate::{
cursor::{
protocol::proto::agent::v1 as pb,
tools::{
codec, edit,
tool_call_result::{self as tool_result, ToolCompletion},
},
},
model::ToolCall,
Error, Result,
};
use super::server_interaction;
pub(crate) fn edit_path_partial(call: &ToolCall, path: &str) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::PartialToolCall(
pb::PartialToolCallUpdate {
call_id: call.call_id.clone(),
tool_call: Some(pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call.call_id.clone()),
started_at_ms: None,
completed_at_ms: None,
tool: Some(pb::tool_call::Tool::EditToolCall(pb::EditToolCall {
args: Some(pb::EditArgs {
path: path.into(),
stream_content: None,
}),
result: None,
})),
}),
args_text_delta: String::new(),
model_call_id: call.model_call_id.clone(),
},
))
}
pub(crate) fn edit_content_delta(call: &ToolCall, content: String) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ToolCallDelta(Box::new(
pb::ToolCallDeltaUpdate {
call_id: call.call_id.clone(),
tool_call_delta: Some(Box::new(pb::ToolCallDelta {
delta: Some(pb::tool_call_delta::Delta::EditToolCallDelta(
pb::EditToolCallDelta {
stream_content_delta: content,
},
)),
})),
model_call_id: call.model_call_id.clone(),
},
)))
}
pub(crate) fn create_plan_partial(
call: &ToolCall,
name: &str,
plan: &str,
overview: &str,
) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::PartialToolCall(
pb::PartialToolCallUpdate {
call_id: call.call_id.clone(),
tool_call: Some(pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call.call_id.clone()),
started_at_ms: None,
completed_at_ms: None,
tool: Some(pb::tool_call::Tool::CreatePlanToolCall(
pb::CreatePlanToolCall {
args: Some(pb::CreatePlanArgs {
plan: plan.into(),
todos: Vec::new(),
overview: overview.into(),
name: name.into(),
is_project: false,
phases: Vec::new(),
}),
result: None,
},
)),
}),
args_text_delta: String::new(),
model_call_id: call.model_call_id.clone(),
},
))
}
pub fn tool_started(
call: &ToolCall,
dynamic_mcp: Option<&pb::McpToolDefinition>,
) -> Result<pb::AgentServerMessage> {
let tool_call = match dynamic_mcp {
Some(definition) => render_dynamic_mcp(call, definition, false),
None => render_tool_call(call, false)?,
};
Ok(server_interaction(
pb::interaction_update::Message::ToolCallStarted(pb::ToolCallStartedUpdate {
call_id: call.call_id.clone(),
tool_call: Some(tool_call),
model_call_id: call.model_call_id.clone(),
}),
))
}
pub fn dynamic_mcp_placeholder(definition: &pb::McpToolDefinition, call_id: &str) -> pb::ToolCall {
dynamic_mcp_tool_call(call_id, None, definition, false, false)
}
pub fn render_dynamic_mcp(
call: &ToolCall,
definition: &pb::McpToolDefinition,
completed: bool,
) -> pb::ToolCall {
dynamic_mcp_tool_call(
&call.call_id,
Some(&call.arguments),
definition,
true,
completed,
)
}
fn dynamic_mcp_tool_call(
call_id: &str,
arguments: Option<&Value>,
definition: &pb::McpToolDefinition,
started: bool,
completed: bool,
) -> pb::ToolCall {
let timestamp = now_ms();
pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call_id.into()),
started_at_ms: started.then_some(timestamp),
completed_at_ms: completed.then_some(timestamp),
tool: Some(pb::tool_call::Tool::McpToolCall(pb::McpToolCall {
args: Some(pb::McpArgs {
name: definition.name.clone(),
args: arguments
.and_then(Value::as_object)
.map(codec::json_object_to_prost)
.unwrap_or_default(),
tool_call_id: call_id.into(),
provider_identifier: definition.provider_identifier.clone(),
tool_name: definition.tool_name.clone(),
..Default::default()
}),
result: None,
description: Some(definition.description.clone()),
})),
}
}
pub fn tool_completed(call: &ToolCall, completion: &ToolCompletion) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::ToolCallCompleted(
pb::ToolCallCompletedUpdate {
call_id: call.call_id.clone(),
tool_call: Some(completion.tool_call().clone()),
model_call_id: call.model_call_id.clone(),
},
))
}
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
use pb::tool_call::Tool;
let tool = match normalized(name).as_str() {
"shell" => Tool::ShellToolCall(pb::ShellToolCall::default()),
"delete" => Tool::DeleteToolCall(pb::DeleteToolCall::default()),
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
"read" => Tool::ReadToolCall(pb::ReadToolCall::default()),
"todowrite" => Tool::UpdateTodosToolCall(pb::UpdateTodosToolCall::default()),
"strreplace" | "editnotebook" | "write" => Tool::EditToolCall(pb::EditToolCall::default()),
"readlints" => Tool::ReadLintsToolCall(pb::ReadLintsToolCall::default()),
"callmcptool" | "semblesearch" | "semblefindrelated" => {
Tool::McpToolCall(pb::McpToolCall::default())
}
"createplan" => Tool::CreatePlanToolCall(pb::CreatePlanToolCall::default()),
"websearch" => Tool::WebSearchToolCall(pb::WebSearchToolCall::default()),
"task" => Tool::TaskToolCall(pb::TaskToolCall::default()),
"fetchmcpresource" => Tool::ReadMcpResourceToolCall(pb::ReadMcpResourceToolCall::default()),
"askquestion" => Tool::AskQuestionToolCall(pb::AskQuestionToolCall::default()),
"webfetch" => Tool::WebFetchToolCall(pb::WebFetchToolCall::default()),
"switchmode" => Tool::SwitchModeToolCall(pb::SwitchModeToolCall::default()),
"generateimage" => Tool::GenerateImageToolCall(pb::GenerateImageToolCall::default()),
"updatecurrentstep" => {
Tool::CommunicateUpdateToolCall(pb::CommunicateUpdateToolCall::default())
}
"getmcptools" => Tool::GetMcpToolsToolCall(pb::GetMcpToolsToolCall::default()),
_ => return Err(Error::Protocol(format!("unsupported tool: {name}"))),
};
Ok(pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call_id.into()),
started_at_ms: None,
completed_at_ms: None,
tool: Some(tool),
})
}
pub fn render_tool_call(call: &ToolCall, completed: bool) -> Result<pb::ToolCall> {
if is_mcp_auth(call) {
let server_identifier = call
.arguments
.get("server")
.and_then(Value::as_str)
.filter(|server| !server.is_empty())
.ok_or_else(|| Error::Protocol("CallMcpTool mcp_auth is missing server".into()))?;
let timestamp = now_ms();
return Ok(pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call.call_id.clone()),
started_at_ms: Some(timestamp),
completed_at_ms: completed.then_some(timestamp),
tool: Some(pb::tool_call::Tool::McpAuthToolCall(pb::McpAuthToolCall {
args: Some(pb::McpAuthArgs {
server_identifier: server_identifier.into(),
tool_call_id: call.call_id.clone(),
}),
result: None,
})),
});
}
let mut output = tool_placeholder(&call.name, &call.call_id)?;
let timestamp = now_ms();
output.started_at_ms = Some(timestamp);
if completed {
output.completed_at_ms = Some(timestamp);
}
let string = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.unwrap_or_default()
.to_string()
};
let optional = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
};
match output.tool.as_mut() {
Some(pb::tool_call::Tool::ShellToolCall(tool)) => {
tool.description = optional("description");
tool.args = Some(pb::ShellArgs {
command: string("command"),
working_directory: optional("working_directory").unwrap_or_default(),
description: optional("description"),
tool_call_id: call.call_id.clone(),
..Default::default()
})
}
Some(pb::tool_call::Tool::DeleteToolCall(tool)) => {
tool.args = Some(pb::DeleteArgs {
path: string("path"),
tool_call_id: call.call_id.clone(),
})
}
Some(pb::tool_call::Tool::GlobToolCall(tool)) => {
tool.args = Some(pb::GlobToolArgs {
target_directory: optional("target_directory"),
glob_pattern: string("glob_pattern"),
})
}
Some(pb::tool_call::Tool::GrepToolCall(tool)) => {
tool.args = Some(pb::GrepArgs {
pattern: string("pattern"),
path: optional("path"),
glob: optional("glob"),
output_mode: optional("output_mode"),
tool_call_id: call.call_id.clone(),
..Default::default()
})
}
Some(pb::tool_call::Tool::ReadToolCall(tool)) => {
tool.args = Some(pb::ReadToolArgs {
path: string("path"),
offset: call
.arguments
.get("offset")
.and_then(Value::as_i64)
.map(|value| value as i32),
limit: call
.arguments
.get("limit")
.and_then(Value::as_i64)
.map(|value| value as i32),
include_line_numbers: call
.arguments
.get("include_line_numbers")
.and_then(Value::as_bool),
})
}
Some(pb::tool_call::Tool::UpdateTodosToolCall(tool)) => {
tool.args = Some(pb::UpdateTodosArgs {
todos: tool_result::todo_items(&call.arguments),
merge: call
.arguments
.get("merge")
.and_then(Value::as_bool)
.unwrap_or(false),
})
}
Some(pb::tool_call::Tool::EditToolCall(tool)) => {
let stream_content = if normalized(&call.name) == "write" {
optional("contents").unwrap_or_default()
} else {
optional("new_string").unwrap_or_default()
};
tool.args = Some(pb::EditArgs {
path: if normalized(&call.name) == "editnotebook" {
string("target_notebook")
} else {
string("path")
},
stream_content: Some(edit::normalize_newlines(&stream_content)),
})
}
Some(pb::tool_call::Tool::ReadLintsToolCall(tool)) => {
tool.args = Some(pb::ReadLintsToolArgs {
paths: call
.arguments
.get("paths")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_string)
.collect(),
})
}
Some(pb::tool_call::Tool::McpToolCall(tool)) => {
tool.description = optional("description");
if let Some(tool_name) = semble_tool_name(&call.name) {
let mut arguments = call.arguments.as_object().cloned().unwrap_or_default();
arguments.remove("description");
tool.args = Some(pb::McpArgs {
name: tool_name.into(),
args: codec::json_object_to_prost(&arguments),
tool_call_id: call.call_id.clone(),
provider_identifier: "builtin-semble".into(),
tool_name: tool_name.into(),
server_identifier: "builtin-semble".into(),
..Default::default()
});
} else {
tool.args = Some(pb::McpArgs {
name: optional("toolName").unwrap_or_default(),
args: call
.arguments
.get("arguments")
.and_then(Value::as_object)
.map(codec::json_object_to_prost)
.unwrap_or_default(),
tool_call_id: call.call_id.clone(),
tool_name: optional("toolName").unwrap_or_default(),
server_identifier: string("server"),
..Default::default()
});
}
}
Some(pb::tool_call::Tool::CreatePlanToolCall(tool)) => {
tool.args = Some(pb::CreatePlanArgs {
plan: string("plan"),
todos: tool_result::todo_items(&call.arguments),
overview: string("overview"),
name: string("name"),
is_project: false,
phases: Vec::new(),
})
}
Some(pb::tool_call::Tool::WebSearchToolCall(tool)) => {
tool.args = Some(pb::WebSearchArgs {
search_term: string("search_term"),
tool_call_id: call.call_id.clone(),
})
}
Some(pb::tool_call::Tool::TaskToolCall(tool)) => {
tool.args = Some(pb::TaskArgs {
description: string("description"),
prompt: string("prompt"),
subagent_type: Some(subagent_type(&string("subagent_type"))),
model: optional("model"),
resume: optional("resume"),
agent_id: None,
attachments: call
.arguments
.get("file_attachments")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_string)
.collect(),
mode: 0,
responding_to_message_ids: Vec::new(),
environment: execution_environment(optional("environment").as_deref()),
machine: None,
})
}
Some(pb::tool_call::Tool::ReadMcpResourceToolCall(tool)) => {
tool.args = Some(pb::ReadMcpResourceExecArgs {
server: string("server"),
uri: string("uri"),
download_path: optional("downloadPath"),
tool_call_id: call.call_id.clone(),
smart_mode_approval: None,
})
}
Some(pb::tool_call::Tool::WebFetchToolCall(tool)) => {
tool.args = Some(pb::WebFetchArgs {
url: string("url"),
tool_call_id: call.call_id.clone(),
})
}
Some(pb::tool_call::Tool::SwitchModeToolCall(tool)) => {
tool.args = Some(pb::SwitchModeArgs {
target_mode_id: string("target_mode_id"),
explanation: optional("explanation"),
tool_call_id: call.call_id.clone(),
})
}
Some(pb::tool_call::Tool::GenerateImageToolCall(tool)) => {
tool.args = Some(pb::GenerateImageArgs {
description: string("description"),
file_path: optional("filename"),
reference_image_paths: call
.arguments
.get("reference_image_paths")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_string)
.collect(),
aspect_ratio: optional("aspect_ratio"),
})
}
Some(pb::tool_call::Tool::CommunicateUpdateToolCall(tool)) => {
tool.args = Some(pb::CommunicateUpdateArgs {
current_step: optional("current_step"),
final_summary: optional("final_summary"),
completed_subtitle: optional("completed_subtitle"),
})
}
Some(pb::tool_call::Tool::WriteShellStdinToolCall(tool)) => {
tool.args = Some(pb::WriteShellStdinArgs {
shell_id: call
.arguments
.get("shell_id")
.and_then(Value::as_u64)
.unwrap_or_default() as u32,
chars: string("chars"),
})
}
Some(pb::tool_call::Tool::GetMcpToolsToolCall(tool)) => {
tool.args = Some(pb::GetMcpToolsArgs {
server: optional("server"),
tool_name: optional("toolName"),
pattern: optional("pattern"),
tool_call_id: call.call_id.clone(),
})
}
_ => {}
}
Ok(output)
}
fn is_mcp_auth(call: &ToolCall) -> bool {
normalized(&call.name) == "callmcptool"
&& call
.arguments
.get("toolName")
.and_then(Value::as_str)
.is_some_and(|tool| normalized(tool) == "mcpauth")
}
fn subagent_type(name: &str) -> pb::SubagentType {
use pb::subagent_type::Type;
let r#type = match name.to_ascii_lowercase().as_str() {
"" | "generalpurpose" => Type::Unspecified(pb::SubagentTypeUnspecified {}),
"explore" => Type::Explore(pb::SubagentTypeExplore {}),
"browser-use" | "browseruse" => Type::BrowserUse(pb::SubagentTypeBrowserUse {}),
"shell" => Type::Shell(pb::SubagentTypeShell {}),
"bash" => Type::Bash(pb::SubagentTypeBash {}),
"debug" => Type::Debug(pb::SubagentTypeDebug {}),
"cursor-guide" | "cursorguide" => Type::CursorGuide(pb::SubagentTypeCursorGuide {}),
"computer-use" | "computeruse" => Type::ComputerUse(pb::SubagentTypeComputerUse {}),
_ => Type::Custom(pb::SubagentTypeCustom { name: name.into() }),
};
pb::SubagentType {
r#type: Some(r#type),
}
}
fn execution_environment(value: Option<&str>) -> i32 {
match value {
Some("cloud") => pb::SubagentExecutionEnvironment::Cloud as i32,
Some("local") | None => pb::SubagentExecutionEnvironment::Local as i32,
Some(_) => pb::SubagentExecutionEnvironment::Unspecified as i32,
}
}
fn normalized(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
fn semble_tool_name(name: &str) -> Option<&'static str> {
match normalized(name).as_str() {
"semblesearch" => Some("search"),
"semblefindrelated" => Some("find_related"),
_ => None,
}
}
fn now_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
+522
View File
@@ -0,0 +1,522 @@
//! Encodes Tool execution requests sent to Cursor.
use serde_json::{Map, Value};
use crate::{
cursor::{
protocol::proto::agent::v1 as pb,
tools::{
edit::{self, EditWrite},
runtime::{ExecContext, McpRoute},
},
},
model::ToolCall,
Error, Result,
};
pub fn request(id: u32, call: &ToolCall, context: &ExecContext) -> Result<pb::AgentServerMessage> {
use pb::exec_server_message::Message;
let string = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
.ok_or_else(|| Error::Protocol(format!("{} is missing {name}", call.name)))
};
let optional_string = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_str)
.map(str::to_string)
};
let int = |name: &str| {
call.arguments
.get(name)
.and_then(Value::as_i64)
.map(|v| v as i32)
};
let message = match normalize(&call.name).as_str() {
"shell" => {
let command = string("command")?;
let (simple_commands, parsing_result) = shell_command_metadata(&command);
Message::ShellStreamArgs(pb::ShellArgs {
command,
working_directory: optional_string("working_directory").unwrap_or_default(),
timeout: shell_timeout(call)?,
tool_call_id: call.call_id.clone(),
simple_commands,
parsing_result,
file_output_threshold_bytes: Some(40_000),
timeout_behavior: pb::TimeoutBehavior::Background as i32,
hard_timeout: Some(86_400_000),
description: optional_string("description"),
output_notification: shell_notification(call)?,
smart_mode_approval: smart_mode_approval(
call,
"request_smart_mode_approval",
"smart_mode_block_reason",
)?,
requested_sandbox_policy: shell_sandbox_policy(call),
close_stdin: true,
conversation_id: Some(context.conversation_id.clone()),
admin_command_denylist: context.admin_command_denylist.clone(),
..Default::default()
})
}
"read" => Message::ReadArgs(pb::ReadArgs {
path: string("path")?,
tool_call_id: call.call_id.clone(),
offset: int("offset"),
limit: call
.arguments
.get("limit")
.and_then(Value::as_u64)
.map(|v| v as u32),
encoding_hint: optional_string("encoding_hint"),
}),
"delete" => Message::DeleteArgs(pb::DeleteArgs {
path: string("path")?,
tool_call_id: call.call_id.clone(),
}),
"grep" => Message::GrepArgs(pb::GrepArgs {
pattern: string("pattern")?,
path: optional_string("path"),
glob: optional_string("glob"),
output_mode: optional_string("output_mode"),
context_before: int("-B"),
context_after: int("-A"),
context: int("-C"),
case_insensitive: call.arguments.get("-i").and_then(Value::as_bool),
r#type: optional_string("type"),
head_limit: int("head_limit"),
multiline: call.arguments.get("multiline").and_then(Value::as_bool),
sort: optional_string("sort"),
sort_ascending: call
.arguments
.get("sort_ascending")
.and_then(Value::as_bool),
tool_call_id: call.call_id.clone(),
sandbox_policy: None,
offset: int("offset"),
}),
"glob" => Message::GrepArgs(pb::GrepArgs {
pattern: String::new(),
path: optional_string("target_directory"),
glob: optional_string("glob_pattern"),
output_mode: Some("files_with_matches".into()),
tool_call_id: call.call_id.clone(),
..Default::default()
}),
"readlints" => Message::DiagnosticsArgs(pb::DiagnosticsArgs {
path: call
.arguments
.get("paths")
.and_then(Value::as_array)
.and_then(|paths| paths.first())
.and_then(Value::as_str)
.unwrap_or_default()
.into(),
tool_call_id: call.call_id.clone(),
}),
"task" => Message::SubagentArgs(pb::SubagentArgs {
tool_call_id: call.call_id.clone(),
subagent_type: optional_string("subagent_type").unwrap_or_default(),
model_id: string("model")?,
prompt: string("prompt")?,
readonly: false,
resume_agent_id: optional_string("resume"),
run_in_background: call
.arguments
.get("run_in_background")
.and_then(Value::as_bool),
continuation_config: None,
parent_conversation_id: Some(context.conversation_id.clone()),
interrupt: call.arguments.get("interrupt").and_then(Value::as_bool),
mode: 0,
fork_agent_id: None,
root_parent_conversation_id: Some(context.root_conversation_id.clone()),
selected_context: task_attachments(call),
direct_meta_parent_child_subagent: None,
environment: match optional_string("environment").as_deref() {
Some("cloud") => pb::SubagentExecutionEnvironment::Cloud as i32,
Some("local") | None => pb::SubagentExecutionEnvironment::Local as i32,
Some(value) => {
return Err(Error::Protocol(format!(
"unknown Task environment: {value}"
)))
}
},
cloud_base_branch: optional_string("cloud_base_branch"),
credentials: None,
}),
"fetchmcpresource" => Message::ReadMcpResourceExecArgs(pb::ReadMcpResourceExecArgs {
server: string("server")?,
uri: string("uri")?,
download_path: optional_string("downloadPath"),
tool_call_id: call.call_id.clone(),
smart_mode_approval: smart_mode_approval(
call,
"requestSmartModeApproval",
"smartModeBlockReason",
)?,
}),
other => {
return Err(Error::Protocol(format!(
"tool {other} is not executed through ExecServerMessage"
)))
}
};
let accept_hook_additional_contexts =
if matches!(&message, pb::exec_server_message::Message::SubagentArgs(_)) {
Some(false)
} else {
Some(true)
};
Ok(server_message(
id,
call,
message,
accept_hook_additional_contexts,
))
}
pub(crate) fn edit_read_request(id: u32, call: &ToolCall) -> Result<pb::AgentServerMessage> {
Ok(server_message(
id,
call,
pb::exec_server_message::Message::ReadArgs(pb::ReadArgs {
path: edit::path(call)?,
tool_call_id: call.call_id.clone(),
..Default::default()
}),
Some(true),
))
}
pub(super) fn edit_write_request(
id: u32,
call: &ToolCall,
write: &EditWrite,
) -> Result<pb::AgentServerMessage> {
Ok(server_message(
id,
call,
pb::exec_server_message::Message::WriteArgs(pb::WriteArgs {
path: edit::path(call)?,
file_text: write.after.clone(),
tool_call_id: call.call_id.clone(),
return_file_content_after_write: false,
file_bytes: Vec::new(),
encoding_hint: None,
}),
Some(true),
))
}
fn server_message(
id: u32,
call: &ToolCall,
message: pb::exec_server_message::Message,
accept_hook_additional_contexts: Option<bool>,
) -> pb::AgentServerMessage {
pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ExecServerMessage(
pb::ExecServerMessage {
id,
exec_id: call.call_id.clone(),
span_context: None,
accept_hook_additional_contexts,
message: Some(message),
},
)),
}
}
pub fn mcp_request(
id: u32,
call: &ToolCall,
definition: &pb::McpToolDefinition,
) -> Result<pb::AgentServerMessage> {
let args = call
.arguments
.as_object()
.map(json_object_to_prost)
.unwrap_or_default();
Ok(pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ExecServerMessage(
pb::ExecServerMessage {
id,
exec_id: call.call_id.clone(),
span_context: None,
accept_hook_additional_contexts: None,
message: Some(pb::exec_server_message::Message::McpArgs(pb::McpArgs {
name: definition.name.clone(),
args,
tool_call_id: call.call_id.clone(),
provider_identifier: definition.provider_identifier.clone(),
tool_name: definition.tool_name.clone(),
smart_mode_approval: None,
smart_mode_approval_only: false,
skip_approval: false,
server_identifier: String::new(),
})),
},
)),
})
}
pub(crate) fn mcp_meta_request(
id: u32,
call: &ToolCall,
server_identifier: &str,
route: &McpRoute,
) -> Result<pb::AgentServerMessage> {
if route.name.is_empty() || route.provider_identifier.is_empty() || route.tool_name.is_empty() {
return Err(Error::Protocol(format!(
"MCP definition for {server_identifier} is incomplete"
)));
}
let requested_tool = call
.arguments
.get("toolName")
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol("CallMcpTool is missing toolName".into()))?;
if requested_tool != route.tool_name {
return Err(Error::Protocol(format!(
"MCP definition mismatch: requested {requested_tool}, resolved {}",
route.tool_name
)));
}
let args = call
.arguments
.get("arguments")
.and_then(Value::as_object)
.map(json_object_to_prost)
.unwrap_or_default();
Ok(server_message(
id,
call,
pb::exec_server_message::Message::McpArgs(pb::McpArgs {
name: route.name.clone(),
args,
tool_call_id: call.call_id.clone(),
provider_identifier: route.provider_identifier.clone(),
tool_name: route.tool_name.clone(),
smart_mode_approval: smart_mode_approval(
call,
"requestSmartModeApproval",
"smartModeBlockReason",
)?,
smart_mode_approval_only: false,
skip_approval: false,
server_identifier: server_identifier.into(),
}),
Some(true),
))
}
pub fn mcp_state_request(id: u32, call: &ToolCall) -> pb::AgentServerMessage {
let server_identifiers = call
.arguments
.get("server")
.and_then(Value::as_str)
.map(|server| vec![server.into()])
.unwrap_or_default();
server_message(
id,
call,
pb::exec_server_message::Message::McpStateExecArgs(pb::McpStateExecArgs {
server_identifiers,
kick_only: false,
}),
Some(false),
)
}
pub fn abort(id: u32) -> pb::AgentServerMessage {
pb::AgentServerMessage {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ExecServerControlMessage(
pb::ExecServerControlMessage {
message: Some(pb::exec_server_control_message::Message::Abort(
pb::ExecServerAbort { id },
)),
},
)),
}
}
fn shell_sandbox_policy(call: &ToolCall) -> Option<pb::SandboxPolicy> {
let permissions = call.arguments.get("required_permissions")?.as_array()?;
let perms: Vec<&str> = permissions.iter().filter_map(Value::as_str).collect();
if perms.contains(&"all") {
Some(pb::SandboxPolicy {
r#type: pb::sandbox_policy::Type::InsecureNone as i32,
network_access: Some(true),
..Default::default()
})
} else if perms.contains(&"full_network") {
Some(pb::SandboxPolicy {
r#type: pb::sandbox_policy::Type::WorkspaceReadwrite as i32,
network_access: Some(true),
..Default::default()
})
} else {
None
}
}
fn shell_command_metadata(command: &str) -> (Vec<String>, Option<pb::ShellCommandParsingResult>) {
let command = command.trim();
let mut parts = command.split_whitespace();
let Some(name) = parts.next() else {
return (Vec::new(), None);
};
let args = parts
.map(
|value| pb::shell_command_parsing_result::ExecutableCommandArg {
r#type: "word".into(),
value: value.into(),
},
)
.collect();
(
vec![command.into()],
Some(pb::ShellCommandParsingResult {
executable_commands: vec![pb::shell_command_parsing_result::ExecutableCommand {
name: name.into(),
args,
full_text: command.into(),
}],
..Default::default()
}),
)
}
fn shell_timeout(call: &ToolCall) -> Result<i32> {
let value = call
.arguments
.get("block_until_ms")
.map(|value| {
value
.as_i64()
.ok_or_else(|| Error::Protocol("Shell block_until_ms must be an integer".into()))
})
.transpose()?
.unwrap_or(30_000);
i32::try_from(value)
.ok()
.filter(|value| *value >= 0)
.ok_or_else(|| Error::Protocol("Shell block_until_ms is out of range".into()))
}
fn smart_mode_approval(
call: &ToolCall,
request_field: &str,
reason_field: &str,
) -> Result<Option<pb::SmartModeApproval>> {
if !call
.arguments
.get(request_field)
.and_then(Value::as_bool)
.unwrap_or(false)
{
return Ok(None);
}
let reason = call
.arguments
.get(reason_field)
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol(format!("{} requires {reason_field}", call.name)))?;
Ok(Some(pb::SmartModeApproval {
request_id: call.call_id.clone(),
reason: reason.to_string(),
}))
}
fn shell_notification(call: &ToolCall) -> Result<Option<pb::ShellOutputNotificationConfig>> {
let Some(value) = call.arguments.get("notify_on_output") else {
return Ok(None);
};
let object = value
.as_object()
.ok_or_else(|| Error::Protocol("Shell notify_on_output must be an object".into()))?;
let required = |field: &str| {
object
.get(field)
.and_then(Value::as_str)
.map(str::to_string)
.ok_or_else(|| Error::Protocol(format!("Shell notify_on_output is missing {field}")))
};
Ok(Some(pb::ShellOutputNotificationConfig {
pattern: required("pattern")?,
reason: required("reason")?,
debounce: object.get("debounce_ms").and_then(Value::as_f64),
notification_limit: None,
}))
}
fn task_attachments(call: &ToolCall) -> Option<pb::SelectedContext> {
let paths = call.arguments.get("file_attachments")?.as_array()?;
let mut context = pb::SelectedContext::default();
for path in paths.iter().filter_map(Value::as_str) {
let extension = std::path::Path::new(path)
.extension()
.and_then(std::ffi::OsStr::to_str)
.unwrap_or_default()
.to_ascii_lowercase();
if matches!(extension.as_str(), "mp4" | "mov" | "webm" | "mkv") {
context.selected_videos.push(pb::SelectedVideo {
path: path.into(),
filename: std::path::Path::new(path)
.file_name()
.and_then(std::ffi::OsStr::to_str)
.unwrap_or_default()
.into(),
materialize_to_filesystem: true,
..Default::default()
});
} else {
context.selected_images.push(pb::SelectedImage {
path: path.into(),
..Default::default()
});
}
}
Some(context)
}
fn normalize(value: &str) -> String {
value
.chars()
.filter(|c| c.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
pub(crate) fn json_object_to_prost(
value: &Map<String, Value>,
) -> std::collections::HashMap<String, prost_types::Value> {
value
.iter()
.map(|(key, value)| (key.clone(), prost_value(value)))
.collect()
}
fn prost_value(value: &Value) -> prost_types::Value {
use prost_types::{value::Kind, ListValue, Struct, Value as ProstValue};
let kind = match value {
Value::Null => Kind::NullValue(0),
Value::Bool(v) => Kind::BoolValue(*v),
Value::Number(v) => Kind::NumberValue(v.as_f64().unwrap_or_default()),
Value::String(v) => Kind::StringValue(v.clone()),
Value::Array(v) => Kind::ListValue(ListValue {
values: v.iter().map(prost_value).collect(),
}),
Value::Object(v) => Kind::StructValue(Struct {
fields: json_object_to_prost(v).into_iter().collect(),
}),
};
ProstValue { kind: Some(kind) }
}
+354
View File
@@ -0,0 +1,354 @@
//! Decodes Tool execution responses received from Cursor.
use crate::{
cursor::{
protocol::{events, proto::agent::v1 as pb},
tools::{
edit,
runtime::{CursorToolRuntime, ExecStage, PendingExec},
tool_call_result::{self as result, ToolCompletion},
},
},
model::ToolCall,
Error, Result,
};
use super::request::edit_write_request;
pub enum ClientExecEvent {
Delta(Box<pb::AgentServerMessage>),
Message(Box<pb::AgentServerMessage>),
Completed(Box<ToolCompletion>),
Pending,
}
pub async fn client_event(
message: &pb::ExecClientMessage,
pending: &CursorToolRuntime,
) -> Result<ClientExecEvent> {
if pending.is_interrupted(message.id).await {
if message.message.as_ref().is_some_and(is_terminal) {
pending.discard_exec(message.id).await;
}
return Ok(ClientExecEvent::Pending);
}
let call = match pending.exec_call(message.id).await {
Some(call) => call,
None if pending.completed_call(message.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal ExecClientMessage id: {}",
message.id
)))
}
None => {
return Err(Error::Protocol(format!(
"unknown ExecClientMessage id: {}",
message.id
)))
}
};
let Some(wire_result) = &message.message else {
return Ok(ClientExecEvent::Pending);
};
let pb::exec_client_message::Message::ShellStream(stream) = wire_result else {
let entry = take(message.id, pending).await?;
return match entry.stage {
ExecStage::EditRead => advance_edit(entry, wire_result, pending).await,
ExecStage::Direct | ExecStage::DynamicMcp(_) | ExecStage::EditWrite(_) => {
completed(entry, wire_result.clone())
}
};
};
use pb::shell_stream::Event;
let event = match &stream.event {
Some(Event::Stdout(stdout)) => {
if pending.append_stdout(message.id, &stdout.data).await {
ClientExecEvent::Delta(Box::new(shell_delta(&call, true, &stdout.data)))
} else {
ClientExecEvent::Pending
}
}
Some(Event::Stderr(stderr)) => {
if pending.append_stderr(message.id, &stderr.data).await {
ClientExecEvent::Delta(Box::new(shell_delta(&call, false, &stderr.data)))
} else {
ClientExecEvent::Pending
}
}
Some(Event::Start(_)) | Some(Event::HookContext(_)) => ClientExecEvent::Pending,
Some(Event::Exit(exit)) => {
let entry = take(message.id, pending).await?;
let result = shell_exit_result(message, exit, &entry.stdout, &entry.stderr);
completed(entry, pb::exec_client_message::Message::ShellResult(result))?
}
Some(Event::Backgrounded(backgrounded)) => {
let entry = take(message.id, pending).await?;
let result = shell_backgrounded_result(
backgrounded,
&entry.stdout,
&entry.stderr,
&entry.context.terminals_folder,
);
completed(entry, pb::exec_client_message::Message::ShellResult(result))?
}
Some(Event::Rejected(value)) => {
let result = pb::ShellResult {
result: Some(pb::shell_result::Result::Rejected(value.clone())),
..Default::default()
};
complete(
message.id,
pending,
pb::exec_client_message::Message::ShellResult(result),
)
.await?
}
Some(Event::PermissionDenied(value)) => {
let result = pb::ShellResult {
result: Some(pb::shell_result::Result::PermissionDenied(value.clone())),
..Default::default()
};
complete(
message.id,
pending,
pb::exec_client_message::Message::ShellResult(result),
)
.await?
}
Some(Event::SandboxUnsupported(value)) => {
let result = pb::ShellResult {
result: Some(pb::shell_result::Result::SpawnError(pb::ShellSpawnError {
command: value.command.clone(),
working_directory: value.working_directory.clone(),
error: value.reason.clone(),
})),
..Default::default()
};
complete(
message.id,
pending,
pb::exec_client_message::Message::ShellResult(result),
)
.await?
}
None => ClientExecEvent::Pending,
};
Ok(event)
}
pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Option<ToolCompletion>> {
if pending.is_interrupted(id).await {
pending.discard_exec(id).await;
return Ok(None);
}
let Some(entry) = pending.take_exec(id).await else {
return Ok(None);
};
let error = "Cursor Exec stream closed before returning a terminal result";
if entry.call.name.eq_ignore_ascii_case("Shell") {
let command = entry
.call
.arguments
.get("command")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string();
let working_directory = entry
.call
.arguments
.get("working_directory")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string();
return Ok(Some(result::from_exec(
entry,
&pb::exec_client_message::Message::ShellResult(pb::ShellResult {
result: Some(pb::shell_result::Result::SpawnError(pb::ShellSpawnError {
command,
working_directory,
error: error.into(),
})),
..Default::default()
}),
)?));
}
let rendered = match &entry.stage {
ExecStage::DynamicMcp(definition) => {
super::render_dynamic_mcp(&entry.call, definition, false)
}
_ => super::render_tool_call(&entry.call, false)?,
};
Ok(Some(ToolCompletion::from_rendered(
&entry.call,
entry.started_at_ms,
error.into(),
true,
rendered,
)?))
}
fn is_terminal(message: &pb::exec_client_message::Message) -> bool {
use pb::{exec_client_message::Message, shell_stream::Event};
match message {
Message::ShellStream(stream) => matches!(
stream.event.as_ref(),
Some(Event::Exit(_))
| Some(Event::Backgrounded(_))
| Some(Event::Rejected(_))
| Some(Event::PermissionDenied(_))
| Some(Event::SandboxUnsupported(_))
),
_ => true,
}
}
async fn advance_edit(
entry: PendingExec,
result: &pb::exec_client_message::Message,
registry: &CursorToolRuntime,
) -> Result<ClientExecEvent> {
let read = match result {
pb::exec_client_message::Message::ReadResult(result)
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
_ => {
return Err(Error::Protocol(format!(
"expected ReadResult for edit tool {}",
entry.call.name
)))
}
};
let write = match edit::after_read(&entry.call, read) {
Ok(write) => write,
Err(error) => {
return Ok(ClientExecEvent::Completed(Box::new(result::edit_failure(
entry, error,
)?)))
}
};
let id = registry
.reserve_edit_write(
&entry.call,
&entry.context,
write.clone(),
entry.started_at_ms,
)
.await?;
Ok(ClientExecEvent::Message(Box::new(edit_write_request(
id,
&entry.call,
&write,
)?)))
}
async fn complete(
id: u32,
pending: &CursorToolRuntime,
result: pb::exec_client_message::Message,
) -> Result<ClientExecEvent> {
completed(take(id, pending).await?, result)
}
async fn take(id: u32, pending: &CursorToolRuntime) -> Result<PendingExec> {
pending
.take_exec(id)
.await
.ok_or_else(|| Error::Protocol(format!("unknown terminal Exec id: {id}")))
}
fn completed(
pending: PendingExec,
result: pb::exec_client_message::Message,
) -> Result<ClientExecEvent> {
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
pending, &result,
)?)))
}
fn shell_exit_result(
message: &pb::ExecClientMessage,
exit: &pb::ShellStreamExit,
stdout: &str,
stderr: &str,
) -> pb::ShellResult {
let result = if exit.code == 0 && !exit.aborted {
pb::shell_result::Result::Success(pb::ShellSuccess {
working_directory: exit.cwd.clone(),
exit_code: exit.code as i32,
stdout: stdout.into(),
stderr: stderr.into(),
interleaved_output: Some(format!("{stdout}{stderr}")),
local_execution_time_ms: exit
.local_execution_time_ms
.or(message.local_execution_time_ms),
..Default::default()
})
} else {
pb::shell_result::Result::Failure(pb::ShellFailure {
working_directory: exit.cwd.clone(),
exit_code: exit.code as i32,
stdout: stdout.into(),
stderr: stderr.into(),
interleaved_output: Some(format!("{stdout}{stderr}")),
abort_reason: exit.abort_reason,
aborted: exit.aborted,
local_execution_time_ms: exit
.local_execution_time_ms
.or(message.local_execution_time_ms),
..Default::default()
})
};
pb::ShellResult {
result: Some(result),
is_background: Some(false),
..Default::default()
}
}
fn shell_backgrounded_result(
backgrounded: &pb::ShellStreamBackgrounded,
stdout: &str,
stderr: &str,
terminals_folder: &str,
) -> pb::ShellResult {
pb::ShellResult {
result: Some(pb::shell_result::Result::Success(pb::ShellSuccess {
command: backgrounded.command.clone(),
working_directory: backgrounded.working_directory.clone(),
stdout: stdout.into(),
stderr: stderr.into(),
shell_id: Some(backgrounded.shell_id),
pid: backgrounded.pid,
ms_to_wait: backgrounded.ms_to_wait,
background_reason: backgrounded.reason,
interleaved_output: Some(format!("{stdout}{stderr}")),
..Default::default()
})),
is_background: Some(true),
terminals_folder: (!terminals_folder.is_empty()).then(|| terminals_folder.into()),
pid: backgrounded.pid,
..Default::default()
}
}
fn shell_delta(call: &ToolCall, stdout: bool, content: &str) -> pb::AgentServerMessage {
let delta = if stdout {
pb::shell_tool_call_delta::Delta::Stdout(pb::ShellToolCallStdoutDelta {
content: content.into(),
})
} else {
pb::shell_tool_call_delta::Delta::Stderr(pb::ShellToolCallStderrDelta {
content: content.into(),
})
};
events::server_interaction(pb::interaction_update::Message::ToolCallDelta(Box::new(
pb::ToolCallDeltaUpdate {
call_id: call.call_id.clone(),
tool_call_delta: Some(Box::new(pb::ToolCallDelta {
delta: Some(pb::tool_call_delta::Delta::ShellToolCallDelta(
pb::ShellToolCallDelta { delta: Some(delta) },
)),
})),
model_call_id: call.model_call_id.clone(),
},
)))
}
+102
View File
@@ -0,0 +1,102 @@
//! Converts unsupported or retired Tool forms into safe Cursor representations.
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{ToolCall, ToolResult},
};
use super::{codec, runtime::now_ms, tool_call_result::ToolCompletion};
// Unknown/retired tools use a generic Cursor MCP card only as a wire/UI
// representation; they are never dispatched to an MCP server.
const COMPAT_PROVIDER: &str = "cursor-byok-compat";
pub(crate) fn placeholder(name: &str, call_id: &str) -> pb::ToolCall {
pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call_id.into()),
started_at_ms: None,
completed_at_ms: None,
tool: Some(pb::tool_call::Tool::McpToolCall(pb::McpToolCall {
args: Some(pb::McpArgs {
name: name.into(),
tool_call_id: call_id.into(),
provider_identifier: COMPAT_PROVIDER.into(),
tool_name: name.into(),
server_identifier: COMPAT_PROVIDER.into(),
..Default::default()
}),
result: None,
description: Some("Unavailable legacy or unsupported tool".into()),
})),
}
}
pub(crate) fn render(call: &ToolCall, completed: bool) -> pb::ToolCall {
let mut output = placeholder(&call.name, &call.call_id);
let timestamp = now_ms();
output.started_at_ms = Some(timestamp);
output.completed_at_ms = completed.then_some(timestamp);
if let Some(pb::tool_call::Tool::McpToolCall(tool)) = output.tool.as_mut() {
if let Some(args) = tool.args.as_mut() {
args.args = call
.arguments
.as_object()
.map(codec::json_object_to_prost)
.unwrap_or_default();
}
}
output
}
pub(crate) fn failure(call: &ToolCall) -> ToolCompletion {
let error = failure_message(&call.name);
let arguments = call
.arguments
.as_object()
.map(codec::json_object_to_prost)
.unwrap_or_default();
ToolCompletion::new(
call,
now_ms(),
ToolResult {
call_id: call.call_id.clone(),
content: error.clone(),
is_error: true,
image: None,
},
pb::tool_call::Tool::McpToolCall(pb::McpToolCall {
args: Some(pb::McpArgs {
name: call.name.clone(),
args: arguments,
tool_call_id: call.call_id.clone(),
provider_identifier: COMPAT_PROVIDER.into(),
tool_name: call.name.clone(),
server_identifier: COMPAT_PROVIDER.into(),
..Default::default()
}),
result: Some(pb::McpToolResult {
result: Some(pb::mcp_tool_result::Result::Error(pb::McpToolError {
error,
read_tool_def_reminder: String::new(),
})),
}),
description: Some("Unavailable legacy or unsupported tool".into()),
}),
)
}
fn failure_message(name: &str) -> String {
if normalized(name) == "awaitshell" {
return "Tool \"AwaitShell\" is no longer available in this Cursor BYOK version. The model emitted a tool name that is not part of the current advertised tool set. Treat the tool call as failed and continue using only tools advertised in the current prompt; for background shell work, use the current Shell/background completion flow.".into();
}
format!(
"Tool \"{name}\" is not available in this Cursor BYOK version. The model emitted a tool name that is not part of the current advertised tool set. Treat the tool call as failed and continue using a tool advertised in the current prompt."
)
}
fn normalized(name: &str) -> String {
name.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
+243
View File
@@ -0,0 +1,243 @@
//! Maintains edit-specific Tool state and projections.
use serde_json::Value;
use similar::{ChangeTag, TextDiff};
use crate::{model::ToolCall, Error, Result};
use crate::cursor::protocol::proto::agent::v1 as pb;
#[derive(Clone, Debug)]
pub(crate) struct EditWrite {
pub before: String,
pub after: String,
}
pub(crate) fn path(call: &ToolCall) -> Result<String> {
let field = if normalized(&call.name) == "editnotebook" {
"target_notebook"
} else {
"path"
};
string(call, field)
}
pub(crate) fn execution_path(call: &ToolCall) -> Result<Option<String>> {
match normalized(&call.name).as_str() {
"write" | "strreplace" | "editnotebook" => path(call).map(Some),
_ => Ok(None),
}
}
pub(crate) fn after_read(
call: &ToolCall,
result: &pb::ReadResult,
) -> std::result::Result<EditWrite, String> {
let before = match result.result.as_ref() {
Some(pb::read_result::Result::Success(success)) => {
if success.truncated {
return Err("cannot edit a truncated Read result".into());
}
match success.output.as_ref() {
Some(pb::read_success::Output::Content(content)) => normalize_newlines(content),
Some(pb::read_success::Output::Data(_)) => {
return Err("cannot edit a binary file".into());
}
None => return Err("Read result has no file content".into()),
}
}
Some(pb::read_result::Result::FileNotFound(_)) if normalized(&call.name) == "write" => {
String::new()
}
Some(pb::read_result::Result::FileNotFound(_)) => {
return Err("file not found".into());
}
Some(pb::read_result::Result::Error(value)) => return Err(value.error.clone()),
Some(pb::read_result::Result::Rejected(value)) => return Err(value.reason.clone()),
Some(pb::read_result::Result::PermissionDenied(_)) => {
return Err("read permission denied".into());
}
Some(pb::read_result::Result::InvalidFile(value)) => {
return Err(value.reason.clone());
}
None => return Err("Read result is empty".into()),
};
let after = match normalized(&call.name).as_str() {
"write" => {
normalize_newlines(&string(call, "contents").map_err(|error| error.to_string())?)
}
"strreplace" => replace_string(call, &before)?,
"editnotebook" => edit_notebook(call, &before)?,
_ => return Err(format!("{} is not an edit tool", call.name)),
};
Ok(EditWrite { before, after })
}
pub(crate) fn success(path: String, write: &EditWrite) -> pb::EditResult {
let diff = TextDiff::from_lines(&write.before, &write.after);
let (mut added, mut removed) = (0, 0);
for change in diff.iter_all_changes() {
match change.tag() {
ChangeTag::Delete => removed += 1,
ChangeTag::Insert => added += 1,
ChangeTag::Equal => {}
}
}
pb::EditResult {
result: Some(pb::edit_result::Result::Success(pb::EditSuccess {
path,
lines_added: Some(added),
lines_removed: Some(removed),
diff_string: Some(diff.unified_diff().to_string()),
before_full_file_content: Some(write.before.clone()),
after_full_file_content: write.after.clone(),
message: None,
})),
}
}
pub(crate) fn failure(path: String, error: impl Into<String>) -> pb::EditResult {
let error = error.into();
pb::EditResult {
result: Some(pb::edit_result::Result::Error(pb::EditError {
path,
error: error.clone(),
model_visible_error: Some(error),
})),
}
}
pub(crate) fn normalize_newlines(value: &str) -> String {
let normalized = value.replace("\r\n", "\n");
normalized.replace('\r', "\n")
}
fn replace_string(call: &ToolCall, before: &str) -> std::result::Result<String, String> {
let old = normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
let new = normalize_newlines(&string(call, "new_string").map_err(|error| error.to_string())?);
if old.is_empty() {
return Err("old_string must not be empty".into());
}
let occurrences = before.match_indices(&old).count();
let replace_all = call
.arguments
.get("replace_all")
.and_then(Value::as_bool)
.unwrap_or(false);
match (replace_all, occurrences) {
(_, 0) => Err("old_string was not found".into()),
(false, 1) => Ok(before.replacen(&old, &new, 1)),
(false, count) => Err(format!(
"old_string is not unique; found {count} occurrences"
)),
(true, _) => Ok(before.replace(&old, &new)),
}
}
fn edit_notebook(call: &ToolCall, before: &str) -> std::result::Result<String, String> {
let mut notebook: Value =
serde_json::from_str(before).map_err(|error| format!("invalid notebook JSON: {error}"))?;
let cells = notebook
.get_mut("cells")
.and_then(Value::as_array_mut)
.ok_or_else(|| "notebook has no cells array".to_string())?;
let index = call
.arguments
.get("cell_idx")
.and_then(Value::as_u64)
.and_then(|value| usize::try_from(value).ok())
.ok_or_else(|| "EditNotebook is missing cell_idx".to_string())?;
let new = normalize_newlines(&string(call, "new_string").map_err(|error| error.to_string())?);
if call
.arguments
.get("is_new_cell")
.and_then(Value::as_bool)
.unwrap_or(false)
{
if index > cells.len() {
return Err(format!("cell_idx {index} is past the end of the notebook"));
}
let language = string(call, "cell_language").map_err(|error| error.to_string())?;
let cell_type = if language == "markdown" || language == "raw" {
language.as_str()
} else {
"code"
};
let mut cell = serde_json::json!({
"cell_type": cell_type,
"metadata": {},
"source": source_lines(&new),
});
if cell_type == "code" {
cell["execution_count"] = Value::Null;
cell["outputs"] = Value::Array(Vec::new());
}
cells.insert(index, cell);
} else {
let cell = cells
.get_mut(index)
.ok_or_else(|| format!("cell_idx {index} does not exist"))?;
let source = cell
.get("source")
.map(notebook_source)
.transpose()?
.unwrap_or_default();
let old =
normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
let occurrences = source.match_indices(&old).count();
let edited = match occurrences {
0 => return Err("old_string was not found in the notebook cell".into()),
1 => source.replacen(&old, &new, 1),
count => {
return Err(format!(
"old_string is not unique in the notebook cell; found {count} occurrences"
))
}
};
cell["source"] = Value::Array(source_lines(&edited));
}
serde_json::to_string_pretty(&notebook)
.map(|value| format!("{value}\n"))
.map_err(|error| error.to_string())
}
fn notebook_source(value: &Value) -> std::result::Result<String, String> {
match value {
Value::String(value) => Ok(normalize_newlines(value)),
Value::Array(lines) => lines
.iter()
.map(|line| {
line.as_str()
.ok_or_else(|| "notebook cell source contains a non-string".to_string())
})
.collect::<std::result::Result<Vec<_>, _>>()
.map(|lines| normalize_newlines(&lines.concat())),
_ => Err("notebook cell source is not text".into()),
}
}
fn source_lines(value: &str) -> Vec<Value> {
if value.is_empty() {
Vec::new()
} else {
value
.split_inclusive('\n')
.map(|line| Value::String(line.to_string()))
.collect()
}
}
fn string(call: &ToolCall, field: &str) -> Result<String> {
call.arguments
.get(field)
.and_then(Value::as_str)
.map(str::to_owned)
.ok_or_else(|| Error::Protocol(format!("{} is missing {field}", call.name)))
}
fn normalized(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
+255
View File
@@ -0,0 +1,255 @@
//! Exposes the extensible Cursor Tool system.
use std::{
collections::{BTreeMap, HashSet},
sync::Arc,
};
use tokio::sync::Mutex;
pub mod codec;
pub(crate) mod compat;
pub(crate) mod edit;
pub(crate) mod registry;
pub mod runtime;
mod schedule;
pub(crate) mod stream;
mod tool_call_dispatch;
pub(crate) mod tool_call_result;
use crate::{
model::{CanonicalMessage, MessageContent, Role, ToolCall},
search::{WebFetch, WebSearch},
store::Store,
Error, Result,
};
use self::schedule::{DeferredEdit, EditSchedule};
use self::tool_call_result::{ToolCompletion, ToolResultSender};
use super::protocol::proto::agent::v1 as pb;
use runtime::{CursorToolRuntime, ExecContext};
#[derive(Clone)]
pub struct ToolDispatcher {
runtime: CursorToolRuntime,
results: ToolResultSender,
search: WebSearch,
fetch: WebFetch,
store: Option<Store>,
edit_schedule: Arc<Mutex<EditSchedule>>,
}
pub struct DispatchedTool {
pub messages: Vec<pb::AgentServerMessage>,
pub completion: Option<ToolCompletion>,
}
pub struct ToolBatchState<'a> {
pub completed: &'a HashSet<String>,
pub started: &'a HashSet<String>,
pub response_text: &'a str,
pub response_thinking: &'a str,
}
pub enum ClientToolEvent {
Completed(Box<ToolCompletion>),
Pending,
}
impl ToolDispatcher {
pub fn new(runtime: CursorToolRuntime) -> Self {
let (results, _) = tool_call_result::tool_result_channel();
Self {
runtime,
results,
search: WebSearch::built_in(),
fetch: WebFetch::built_in(),
store: None,
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
}
}
pub fn with_results(
runtime: CursorToolRuntime,
results: ToolResultSender,
store: Store,
) -> Self {
Self {
runtime,
results,
search: WebSearch::managed(store.clone()),
fetch: WebFetch::managed(store.clone()),
store: Some(store),
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
}
}
pub async fn start_batch(
&self,
calls: &[ToolCall],
state: ToolBatchState<'_>,
messages: &[CanonicalMessage],
dynamic_mcp: &BTreeMap<String, pb::McpToolDefinition>,
context: &ExecContext,
) -> Result<Vec<DispatchedTool>> {
let first_tool_index = current_turn_step_count(messages)
+ usize::from(!state.response_thinking.is_empty())
+ usize::from(!state.response_text.is_empty())
+ 1;
let mut dispatched = Vec::new();
for (position, call) in calls.iter().enumerate() {
if state.completed.contains(&call.call_id) {
continue;
}
let message_index = first_tool_index + position;
let publish_started = !state.started.contains(&call.call_id);
let edit_path = if dynamic_mcp.contains_key(&call.name) {
None
} else {
edit::execution_path(call)?
};
if let Some(path) = edit_path {
let next = self.edit_schedule.lock().await.start_or_defer(
path,
DeferredEdit {
call: call.clone(),
message_index,
publish_started,
context: context.clone(),
},
);
let Some(next) = next else {
continue;
};
dispatched.push(
self.start(
&next.call,
next.message_index,
next.publish_started,
dynamic_mcp,
&next.context,
)
.await?,
);
continue;
}
dispatched.push(
self.start(call, message_index, publish_started, dynamic_mcp, context)
.await?,
);
}
Ok(dispatched)
}
pub(crate) async fn continue_after(&self, call_id: &str) -> Result<Option<DispatchedTool>> {
let next = self.edit_schedule.lock().await.complete(call_id)?;
let Some(next) = next else {
return Ok(None);
};
self.start(
&next.call,
next.message_index,
next.publish_started,
&BTreeMap::new(),
&next.context,
)
.await
.map(Some)
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
self.edit_schedule.lock().await.clear();
self.runtime.interrupt_for_message().await
}
async fn start(
&self,
call: &ToolCall,
message_index: usize,
publish_started: bool,
dynamic_mcp: &BTreeMap<String, pb::McpToolDefinition>,
context: &ExecContext,
) -> Result<DispatchedTool> {
let call = context.prepare_call(call)?;
let mut messages = if publish_started {
vec![codec::tool_started(&call, dynamic_mcp.get(&call.name))?]
} else {
Vec::new()
};
let started = tool_call_dispatch::start(
&self.runtime,
&self.results,
&call,
message_index,
dynamic_mcp,
context,
self.store.as_ref(),
)
.await?;
messages.extend(started.messages);
Ok(DispatchedTool {
messages,
completion: started.completion,
})
}
pub async fn interaction_response(
&self,
response: &pb::InteractionResponse,
) -> Result<ClientToolEvent> {
if self.runtime.is_interrupted(response.id).await {
return Ok(ClientToolEvent::Pending);
}
let pending = match self.runtime.take_interaction(response.id).await {
Some(pending) => pending,
None if self.runtime.completed_call(response.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal InteractionResponse id: {}",
response.id
)));
}
None => {
return Err(Error::Protocol(format!(
"unknown InteractionResponse id: {}",
response.id
)));
}
};
Ok(
match tool_call_dispatch::resume_interaction(
&self.results,
&self.search,
&self.fetch,
pending,
response,
)
.await?
{
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
ClientToolEvent::Completed(completion)
}
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
},
)
}
}
fn current_turn_step_count(messages: &[CanonicalMessage]) -> usize {
let turn_start = messages
.iter()
.rposition(|message| message.role == Role::User)
.map_or(0, |position| position + 1);
messages[turn_start..]
.iter()
.map(|message| match &message.content {
MessageContent::Assistant {
text,
thinking,
tool_calls,
..
} => {
usize::from(!thinking.is_empty()) + usize::from(!text.is_empty()) + tool_calls.len()
}
_ => 0,
})
.sum()
}
+91
View File
@@ -0,0 +1,91 @@
//! Owns the static Cursor Tool schema catalog and mode-specific selections.
use std::collections::HashMap;
use serde::Deserialize;
use serde_json::Value;
use crate::{model::ToolDefinition, Error, Result};
#[derive(Deserialize)]
struct Manifest {
tools: Vec<ManifestTool>,
}
#[derive(Deserialize)]
#[serde(untagged)]
enum ManifestTool {
Name(String),
Variant { name: String, variant: String },
}
pub(crate) struct ToolRegistry {
tools: HashMap<String, ToolDefinition>,
variants: HashMap<String, ToolDefinition>,
}
impl ToolRegistry {
pub(crate) fn parse(json: &str) -> Result<Self> {
let value: Value = serde_json::from_str(json)?;
let tools = value
.get("tools")
.and_then(Value::as_array)
.ok_or_else(|| Error::Config("tools.json is missing tools".into()))?
.iter()
.map(parse_tool)
.map(|result| result.map(|tool| (tool.name.clone(), tool)))
.collect::<Result<HashMap<_, _>>>()?;
let variants = value
.get("variants")
.and_then(Value::as_object)
.into_iter()
.flat_map(|variants| variants.iter())
.map(|(name, value)| parse_tool(value).map(|tool| (name.clone(), tool)))
.collect::<Result<HashMap<_, _>>>()?;
Ok(Self { tools, variants })
}
pub(crate) fn select_json(&self, manifest: &str) -> Result<Vec<ToolDefinition>> {
let manifest: Manifest = serde_json::from_str(manifest)?;
manifest
.tools
.iter()
.map(|entry| match entry {
ManifestTool::Name(name) => self.tools.get(name).cloned().ok_or_else(|| {
Error::Config(format!("tool manifest references unknown schema: {name}"))
}),
ManifestTool::Variant { name, variant } => self
.variants
.get(&format!("{name}.{variant}"))
.cloned()
.ok_or_else(|| {
Error::Config(format!(
"tool manifest references unknown variant: {name}.{variant}"
))
}),
})
.collect()
}
}
fn parse_tool(tool: &Value) -> Result<ToolDefinition> {
let function = tool
.get("function")
.ok_or_else(|| Error::Config("tool is missing function".into()))?;
Ok(ToolDefinition {
name: function
.get("name")
.and_then(Value::as_str)
.ok_or_else(|| Error::Config("tool is missing name".into()))?
.into(),
description: function
.get("description")
.and_then(Value::as_str)
.ok_or_else(|| Error::Config("tool is missing description".into()))?
.into(),
parameters: function
.get("parameters")
.cloned()
.ok_or_else(|| Error::Config("tool is missing parameters".into()))?,
})
}
+367
View File
@@ -0,0 +1,367 @@
//! Tracks running Tool executions and coordinates cancellation and cleanup.
use std::{
collections::{HashMap, HashSet},
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
};
use tokio::sync::Mutex;
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolCall, Error, Result};
use super::edit::EditWrite;
#[derive(Clone, Default)]
pub struct CursorToolRuntime {
next_id: Arc<AtomicU32>,
execs: Arc<Mutex<HashMap<u32, PendingExec>>>,
interactions: Arc<Mutex<HashMap<u32, PendingInteraction>>>,
completed: Arc<Mutex<HashMap<u32, String>>>,
interrupted: Arc<Mutex<HashSet<u32>>>,
}
pub(crate) struct PendingExec {
pub call: ToolCall,
pub context: ExecContext,
pub started_at_ms: u64,
pub stdout: String,
pub stderr: String,
pub stage: ExecStage,
}
pub(crate) enum ExecStage {
Direct,
DynamicMcp(pb::McpToolDefinition),
EditRead,
EditWrite(EditWrite),
}
#[derive(Clone, Debug, Default)]
pub struct ExecContext {
pub conversation_id: String,
pub root_conversation_id: String,
pub default_subagent_model: String,
pub subagent_model: Option<SubagentModel>,
pub allow_subagents: bool,
pub subagents_disabled: bool,
pub terminals_folder: String,
pub admin_command_denylist: Vec<String>,
pub mcp_routes: HashMap<(String, String), McpRoute>,
}
#[derive(Clone, Debug)]
pub struct McpRoute {
pub name: String,
pub provider_identifier: String,
pub tool_name: String,
pub description: String,
}
#[derive(Clone, Debug)]
pub enum SubagentModel {
Model(String),
Disabled,
}
impl ExecContext {
pub fn task_disabled(&self, call: &ToolCall) -> bool {
if !call.name.eq_ignore_ascii_case("Task") {
return false;
}
self.subagents_disabled || matches!(self.subagent_model, Some(SubagentModel::Disabled))
}
pub fn prepare_call(&self, call: &ToolCall) -> Result<ToolCall> {
if !call.name.eq_ignore_ascii_case("Task") {
return Ok(call.clone());
}
let arguments = call
.arguments
.as_object()
.ok_or_else(|| Error::Protocol("Task arguments must be a JSON object".into()))?;
let subagent_type = arguments
.get("subagent_type")
.and_then(serde_json::Value::as_str)
.unwrap_or("generalPurpose");
if self.task_disabled(call) {
return Ok(call.clone());
}
let model = match &self.subagent_model {
Some(SubagentModel::Model(model)) => model.clone(),
Some(SubagentModel::Disabled) => unreachable!("disabled Task returned above"),
None => arguments
.get("model")
.and_then(serde_json::Value::as_str)
.filter(|model| *model != "inherit")
.unwrap_or(&self.default_subagent_model)
.to_string(),
};
if model.is_empty() {
return Err(Error::Protocol(format!(
"Task subagent type {subagent_type} has no model"
)));
}
let mut prepared = call.clone();
prepared
.arguments
.as_object_mut()
.expect("Task arguments were validated")
.insert("model".into(), serde_json::Value::String(model));
Ok(prepared)
}
}
pub(crate) struct PendingInteraction {
pub call: ToolCall,
pub started_at_ms: u64,
}
impl CursorToolRuntime {
pub(crate) fn next_run(&self) -> Self {
Self {
next_id: self.next_id.clone(),
execs: Arc::new(Mutex::new(HashMap::new())),
interactions: Arc::new(Mutex::new(HashMap::new())),
completed: Arc::new(Mutex::new(HashMap::new())),
interrupted: self.interrupted.clone(),
}
}
pub async fn reserve_exec(&self, call: &ToolCall, context: &ExecContext) -> Result<u32> {
self.reserve_exec_stage(call, context, ExecStage::Direct, None)
.await
}
pub(crate) async fn reserve_dynamic_mcp(
&self,
call: &ToolCall,
context: &ExecContext,
definition: &pb::McpToolDefinition,
) -> Result<u32> {
self.reserve_exec_stage(
call,
context,
ExecStage::DynamicMcp(definition.clone()),
None,
)
.await
}
pub(crate) async fn reserve_edit_read(
&self,
call: &ToolCall,
context: &ExecContext,
) -> Result<u32> {
self.reserve_exec_stage(call, context, ExecStage::EditRead, None)
.await
}
pub(crate) async fn reserve_edit_write(
&self,
call: &ToolCall,
context: &ExecContext,
write: EditWrite,
started_at_ms: u64,
) -> Result<u32> {
self.reserve_exec_stage(
call,
context,
ExecStage::EditWrite(write),
Some(started_at_ms),
)
.await
}
async fn reserve_exec_stage(
&self,
call: &ToolCall,
context: &ExecContext,
stage: ExecStage,
started_at_ms: Option<u64>,
) -> Result<u32> {
let id = self.next_id()?;
self.execs.lock().await.insert(
id,
PendingExec {
call: call.clone(),
context: context.clone(),
started_at_ms: started_at_ms.unwrap_or_else(now_ms),
stdout: String::new(),
stderr: String::new(),
stage,
},
);
Ok(id)
}
pub async fn reserve_interaction(&self, call: &ToolCall) -> Result<u32> {
let id = self.next_id()?;
self.interactions.lock().await.insert(
id,
PendingInteraction {
call: call.clone(),
started_at_ms: now_ms(),
},
);
Ok(id)
}
pub async fn exec_call(&self, id: u32) -> Option<ToolCall> {
self.execs
.lock()
.await
.get(&id)
.map(|entry| entry.call.clone())
}
pub async fn append_stdout(&self, id: u32, data: &str) -> bool {
let mut entries = self.execs.lock().await;
let Some(entry) = entries.get_mut(&id) else {
return false;
};
entry.stdout.push_str(data);
true
}
pub async fn append_stderr(&self, id: u32, data: &str) -> bool {
let mut entries = self.execs.lock().await;
let Some(entry) = entries.get_mut(&id) else {
return false;
};
entry.stderr.push_str(data);
true
}
pub(crate) async fn take_exec(&self, id: u32) -> Option<PendingExec> {
let pending = self.execs.lock().await.remove(&id);
if let Some(pending) = &pending {
self.completed
.lock()
.await
.insert(id, pending.call.call_id.clone());
}
pending
}
pub(crate) async fn take_interaction(&self, id: u32) -> Option<PendingInteraction> {
let pending = self.interactions.lock().await.remove(&id);
if let Some(pending) = &pending {
self.completed
.lock()
.await
.insert(id, pending.call.call_id.clone());
}
pending
}
pub async fn completed_call(&self, id: u32) -> Option<String> {
self.completed.lock().await.get(&id).cloned()
}
pub async fn is_interrupted(&self, id: u32) -> bool {
self.interrupted.lock().await.contains(&id)
}
pub async fn clear_completed(&self) {
self.completed.lock().await.clear();
}
pub async fn discard_exec(&self, id: u32) {
self.execs.lock().await.remove(&id);
}
pub async fn discard_interaction(&self, id: u32) {
self.interactions.lock().await.remove(&id);
}
pub async fn drain_running(&self) -> Vec<u32> {
let mut entries = self.execs.lock().await;
let mut ids = entries.drain().map(|(id, _)| id).collect::<Vec<_>>();
ids.sort_unstable();
self.interactions.lock().await.clear();
self.completed.lock().await.clear();
self.interrupted.lock().await.clear();
ids
}
pub async fn interrupt_for_run_replacement(&self) -> Vec<u32> {
let mut execs = self.execs.lock().await;
let mut abort_ids = execs.keys().copied().collect::<Vec<_>>();
let mut interrupted_ids = abort_ids.clone();
execs.clear();
drop(execs);
let mut interactions = self.interactions.lock().await;
interrupted_ids.extend(interactions.keys().copied());
interactions.clear();
drop(interactions);
self.completed.lock().await.clear();
self.interrupted.lock().await.extend(interrupted_ids);
abort_ids.sort_unstable();
abort_ids
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
let (abort_ids, interrupted_ids) = {
let mut entries = self.execs.lock().await;
let mut abort_ids = Vec::new();
let mut interrupted_ids = Vec::new();
entries.retain(|id, entry| {
interrupted_ids.push(*id);
let keep_running = entry.call.name.eq_ignore_ascii_case("Task");
if !keep_running {
abort_ids.push(*id);
}
keep_running
});
(abort_ids, interrupted_ids)
};
let interaction_ids = {
let mut interactions = self.interactions.lock().await;
let ids = interactions.keys().copied().collect::<Vec<_>>();
interactions.clear();
ids
};
let mut interrupted = self.interrupted.lock().await;
interrupted.extend(interrupted_ids);
interrupted.extend(interaction_ids);
let mut abort_ids = abort_ids;
abort_ids.sort_unstable();
abort_ids
}
pub async fn running_exec_ids(&self) -> Vec<u32> {
let mut ids = self.execs.lock().await.keys().copied().collect::<Vec<_>>();
ids.sort_unstable();
ids
}
pub async fn running_task_exec_id(&self, call_id: &str) -> Option<u32> {
self.execs
.lock()
.await
.iter()
.filter_map(|(id, entry)| {
(entry.call.call_id == call_id && entry.call.name.eq_ignore_ascii_case("Task"))
.then_some(*id)
})
.min()
}
fn next_id(&self) -> Result<u32> {
self.next_id
.fetch_add(1, Ordering::Relaxed)
.checked_add(1)
.ok_or_else(|| Error::Protocol("Cursor message id space exhausted".into()))
}
}
pub(crate) fn now_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
+74
View File
@@ -0,0 +1,74 @@
//! Schedules background Tool work.
use std::collections::{HashMap, VecDeque};
use crate::{model::ToolCall, Error, Result};
use super::runtime::ExecContext;
#[derive(Default)]
pub(super) struct EditSchedule {
paths: HashMap<String, EditPathQueue>,
active_paths: HashMap<String, String>,
}
struct EditPathQueue {
active_call_id: String,
waiting: VecDeque<DeferredEdit>,
}
pub(super) struct DeferredEdit {
pub call: ToolCall,
pub message_index: usize,
pub publish_started: bool,
pub context: ExecContext,
}
impl EditSchedule {
pub fn clear(&mut self) {
self.paths.clear();
self.active_paths.clear();
}
pub fn start_or_defer(&mut self, path: String, edit: DeferredEdit) -> Option<DeferredEdit> {
if let Some(queue) = self.paths.get_mut(&path) {
queue.waiting.push_back(edit);
return None;
}
self.active_paths
.insert(edit.call.call_id.clone(), path.clone());
self.paths.insert(
path,
EditPathQueue {
active_call_id: edit.call.call_id.clone(),
waiting: VecDeque::new(),
},
);
Some(edit)
}
pub fn complete(&mut self, call_id: &str) -> Result<Option<DeferredEdit>> {
let Some(path) = self.active_paths.remove(call_id) else {
return Ok(None);
};
let queue = self.paths.get_mut(&path).ok_or_else(|| {
Error::Protocol(format!("active edit path disappeared for call {call_id}"))
})?;
if queue.active_call_id != call_id {
return Err(Error::Protocol(format!(
"edit path is active for {}, not {call_id}",
queue.active_call_id
)));
}
match queue.waiting.pop_front() {
Some(next) => {
queue.active_call_id = next.call.call_id.clone();
self.active_paths.insert(next.call.call_id.clone(), path);
Ok(Some(next))
}
None => {
self.paths.remove(&path);
Ok(None)
}
}
}
}
+186
View File
@@ -0,0 +1,186 @@
//! Projects streaming Tool arguments to Cursor updates.
use crate::{
cursor::{
protocol::{
json_stream::{JsonStringFields, StringFieldEvent},
proto::agent::v1 as pb,
},
tools::codec as interaction,
},
model::ToolCall,
Result,
};
pub struct ToolCallStream {
presentation: Presentation,
}
enum Presentation {
Plain,
DynamicMcp(pb::McpToolDefinition),
Edit(EditProjection),
CreatePlan(CreatePlanProjection),
}
struct EditProjection {
fields: JsonStringFields,
path_field: &'static str,
content_field: &'static str,
path: String,
content: NewlineStream,
}
#[derive(Default)]
struct CreatePlanProjection {
fields: JsonStringFields,
name: String,
plan: String,
overview: String,
}
impl ToolCallStream {
pub fn new(name: &str, dynamic_mcp: Option<&pb::McpToolDefinition>) -> Self {
let presentation = match dynamic_mcp {
Some(definition) => Presentation::DynamicMcp(definition.clone()),
None => match normalized(name).as_str() {
"write" => Presentation::Edit(EditProjection::new("path", "contents")),
"strreplace" => Presentation::Edit(EditProjection::new("path", "new_string")),
"editnotebook" => {
Presentation::Edit(EditProjection::new("target_notebook", "new_string"))
}
"createplan" => Presentation::CreatePlan(CreatePlanProjection::default()),
_ => Presentation::Plain,
},
};
Self { presentation }
}
pub fn arguments_delta(
&mut self,
call: &ToolCall,
raw_delta: &str,
) -> Result<Vec<pb::AgentServerMessage>> {
match &mut self.presentation {
Presentation::Plain => Ok(vec![interaction::arguments_delta(call, raw_delta)?]),
Presentation::DynamicMcp(definition) => {
Ok(vec![interaction::dynamic_mcp_arguments_delta(
call, raw_delta, definition,
)])
}
Presentation::Edit(edit) => {
let mut messages = Vec::new();
edit.project(call, raw_delta, &mut messages)?;
Ok(messages)
}
Presentation::CreatePlan(plan) => plan.project(call, raw_delta),
}
}
}
impl CreatePlanProjection {
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
let mut completed_field = false;
for event in self.fields.push(raw_delta)? {
match event {
StringFieldEvent::Delta { name, text } => match name.as_str() {
"name" => self.name.push_str(&text),
"plan" => self.plan.push_str(&text),
"overview" => self.overview.push_str(&text),
_ => {}
},
StringFieldEvent::End { name }
if matches!(name.as_str(), "name" | "plan" | "overview") =>
{
completed_field = true
}
_ => {}
}
}
Ok(completed_field
.then(|| interaction::create_plan_partial(call, &self.name, &self.plan, &self.overview))
.into_iter()
.collect())
}
}
impl EditProjection {
fn new(path_field: &'static str, content_field: &'static str) -> Self {
Self {
fields: JsonStringFields::default(),
path_field,
content_field,
path: String::new(),
content: NewlineStream::default(),
}
}
fn project(
&mut self,
call: &ToolCall,
raw_delta: &str,
messages: &mut Vec<pb::AgentServerMessage>,
) -> Result<()> {
for event in self.fields.push(raw_delta)? {
match event {
StringFieldEvent::Delta { name, text } if name == self.path_field => {
self.path.push_str(&text)
}
StringFieldEvent::End { name } if name == self.path_field => {
messages.push(interaction::edit_path_partial(call, &self.path));
}
StringFieldEvent::Delta { name, text } if name == self.content_field => {
let content = self.content.push(&text, false);
if !content.is_empty() {
messages.push(interaction::edit_content_delta(call, content));
}
}
StringFieldEvent::End { name } if name == self.content_field => {
let content = self.content.push("", true);
if !content.is_empty() {
messages.push(interaction::edit_content_delta(call, content));
}
}
_ => {}
}
}
Ok(())
}
}
#[derive(Default)]
struct NewlineStream {
pending_cr: bool,
}
impl NewlineStream {
fn push(&mut self, text: &str, finished: bool) -> String {
let mut output = String::with_capacity(text.len());
for character in text.chars() {
if self.pending_cr {
output.push('\n');
self.pending_cr = false;
if character == '\n' {
continue;
}
}
if character == '\r' {
self.pending_cr = true;
} else {
output.push(character);
}
}
if finished && self.pending_cr {
output.push('\n');
self.pending_cr = false;
}
output
}
}
fn normalized(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
@@ -0,0 +1,22 @@
//! Dispatches edit Tool calls.
//! Hidden read phase for file editing tools.
use crate::{model::ToolCall, Result};
use super::ToolStart;
use crate::cursor::tools::{
codec,
runtime::{CursorToolRuntime, ExecContext},
};
pub(super) async fn start(
runtime: &CursorToolRuntime,
call: &ToolCall,
context: &ExecContext,
) -> Result<ToolStart> {
let id = runtime.reserve_edit_read(call, context).await?;
Ok(ToolStart {
messages: vec![codec::edit_read_request(id, call)?],
completion: None,
})
}
@@ -0,0 +1,73 @@
//! Dispatches command execution Tool calls.
//! Direct Exec and dynamic MCP dispatch.
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolCall, Error, Result};
use super::{normalized, ToolStart};
use crate::cursor::tools::{
codec,
runtime::{CursorToolRuntime, ExecContext},
tool_call_result as result,
};
pub(super) async fn start(
runtime: &CursorToolRuntime,
call: &ToolCall,
context: &ExecContext,
) -> Result<ToolStart> {
let message = match normalized(&call.name).as_str() {
"getmcptools" => {
let id = runtime.reserve_exec(call, context).await?;
codec::mcp_state_request(id, call)
}
"callmcptool" => {
let server = required(call, "server")?;
let tool = required(call, "toolName")?;
let Some(route) = context
.mcp_routes
.get(&(server.to_string(), tool.to_string()))
else {
return Ok(ToolStart {
messages: Vec::new(),
completion: Some(result::mcp_failure(
call,
format!("MCP descriptor not found for {server}/{tool}"),
)?),
});
};
let id = runtime.reserve_exec(call, context).await?;
codec::mcp_meta_request(id, call, server, route)?
}
_ => {
let id = runtime.reserve_exec(call, context).await?;
codec::request(id, call, context)?
}
};
Ok(ToolStart {
messages: vec![message],
completion: None,
})
}
fn required<'a>(call: &'a ToolCall, name: &str) -> Result<&'a str> {
call.arguments
.get(name)
.and_then(serde_json::Value::as_str)
.filter(|value| !value.is_empty())
.ok_or_else(|| Error::Protocol(format!("{} is missing {name}", call.name)))
}
pub(super) async fn start_dynamic(
runtime: &CursorToolRuntime,
call: &ToolCall,
definition: &pb::McpToolDefinition,
context: &ExecContext,
) -> Result<ToolStart> {
let id = runtime
.reserve_dynamic_mcp(call, context, definition)
.await?;
Ok(ToolStart {
messages: vec![codec::mcp_request(id, call, definition)?],
completion: None,
})
}
@@ -0,0 +1,110 @@
//! Dispatches Tool calls that require Cursor user interaction.
//! Interaction query dispatch and approval continuation.
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec as interaction},
model::ToolCall,
search::{WebFetch, WebSearch},
Error, Result,
};
use super::{normalized, InteractionContinuation, ToolStart};
use crate::cursor::tools::{
runtime::{CursorToolRuntime, PendingInteraction},
tool_call_result::{self as result, ToolResultSender},
};
pub(super) async fn start(runtime: &CursorToolRuntime, call: &ToolCall) -> Result<ToolStart> {
let id = runtime.reserve_interaction(call).await?;
Ok(ToolStart {
messages: vec![interaction::tool_query(id, call)?],
completion: None,
})
}
pub(super) async fn resume(
results: &ToolResultSender,
search: &WebSearch,
fetch: &WebFetch,
pending: PendingInteraction,
response: &pb::InteractionResponse,
) -> Result<InteractionContinuation> {
if normalized(&pending.call.name) == "websearch"
&& matches!(
response.result.as_ref(),
Some(pb::interaction_response::Result::WebSearchRequestResponse(
pb::WebSearchRequestResponse {
result: Some(pb::web_search_request_response::Result::Approved(_)),
}
))
)
{
start_web_search(results.clone(), search.clone(), pending)?;
return Ok(InteractionContinuation::Pending);
}
if normalized(&pending.call.name) == "webfetch"
&& matches!(
response.result.as_ref(),
Some(pb::interaction_response::Result::WebFetchRequestResponse(
pb::WebFetchRequestResponse {
result: Some(pb::web_fetch_request_response::Result::Approved(_)),
}
))
)
{
start_web_fetch(results.clone(), fetch.clone(), pending)?;
return Ok(InteractionContinuation::Pending);
}
Ok(InteractionContinuation::Completed(Box::new(
result::from_interaction(pending, response)?,
)))
}
fn start_web_fetch(
results: ToolResultSender,
fetch: WebFetch,
pending: PendingInteraction,
) -> Result<()> {
let url = pending
.call
.arguments
.get("url")
.and_then(serde_json::Value::as_str)
.filter(|url| !url.trim().is_empty())
.ok_or_else(|| Error::Protocol("WebFetch is missing url".into()))?
.to_string();
tokio::spawn(async move {
let outcome = fetch.fetch(&url).await.map_err(|error| error.to_string());
match result::complete_web_fetch(pending, outcome) {
Ok(completion) => results.send(completion),
Err(error) => results.send_error(error),
}
});
Ok(())
}
fn start_web_search(
results: ToolResultSender,
search: WebSearch,
pending: PendingInteraction,
) -> Result<()> {
let query = pending
.call
.arguments
.get("search_term")
.and_then(serde_json::Value::as_str)
.filter(|query| !query.trim().is_empty())
.ok_or_else(|| Error::Protocol("WebSearch is missing search_term".into()))?
.to_string();
tokio::spawn(async move {
let outcome = search
.search(&query)
.await
.map_err(|error| error.to_string());
match result::complete_web_search(pending, outcome) {
Ok(completion) => results.send(completion),
Err(error) => results.send_error(error),
}
});
Ok(())
}
@@ -0,0 +1,21 @@
//! Dispatches server-local Tool calls.
//! Synchronous local tool dispatch.
use crate::{model::ToolCall, Result};
use super::ToolStart;
use crate::cursor::tools::tool_call_result as result;
pub(super) fn start(call: &ToolCall, message_index: usize) -> Result<ToolStart> {
Ok(ToolStart {
messages: Vec::new(),
completion: Some(result::local(call, message_index)?),
})
}
pub(super) fn subagents_disabled(call: &ToolCall) -> Result<ToolStart> {
Ok(ToolStart {
messages: Vec::new(),
completion: Some(result::subagents_disabled(call)?),
})
}
@@ -0,0 +1,156 @@
//! Dispatches Tool calls to their execution adapters.
mod edit;
mod exec;
mod interaction;
mod local;
mod search;
use std::collections::BTreeMap;
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::ToolCall,
search::{WebFetch, WebSearch},
store::Store,
Error, Result,
};
use super::{
compat,
runtime::{CursorToolRuntime, ExecContext, PendingInteraction},
tool_call_result::{ToolCompletion, ToolResultSender},
};
pub(super) struct ToolStart {
pub messages: Vec<pb::AgentServerMessage>,
pub completion: Option<ToolCompletion>,
}
pub(super) enum InteractionContinuation {
Completed(Box<ToolCompletion>),
Pending,
}
pub(super) async fn start(
runtime: &CursorToolRuntime,
results: &ToolResultSender,
call: &ToolCall,
message_index: usize,
dynamic_mcp: &BTreeMap<String, pb::McpToolDefinition>,
context: &ExecContext,
store: Option<&Store>,
) -> Result<ToolStart> {
if let Some(definition) = dynamic_mcp.get(&call.name) {
return exec::start_dynamic(runtime, call, definition, context).await;
}
if is_mcp_auth(call) {
return interaction::start(runtime, call).await;
}
if context.task_disabled(call) {
return local::subagents_disabled(call);
}
let normalized_call = normalize_block_until_ms(call)?;
let call = normalized_call.as_ref().unwrap_or(call);
match normalized(&call.name).as_str() {
"shell" | "bash" | "read" | "delete" | "grep" | "glob" | "readlints" | "task"
| "callmcptool" | "fetchmcpresource" | "getmcptools" => {
exec::start(runtime, call, context).await
}
"write" | "strreplace" | "editnotebook" => edit::start(runtime, call, context).await,
"askquestion" | "websearch" | "webfetch" | "switchmode" | "createplan"
| "generateimage" => interaction::start(runtime, call).await,
"todowrite" | "updatecurrentstep" => local::start(call, message_index),
"semblesearch" | "semblefindrelated" => search::start(results, call, store.cloned()),
_ => Ok(unavailable_tool(call)),
}
}
fn unavailable_tool(call: &ToolCall) -> ToolStart {
ToolStart {
messages: Vec::new(),
completion: Some(compat::failure(call)),
}
}
fn normalize_block_until_ms(call: &ToolCall) -> Result<Option<ToolCall>> {
if !is_shell_tool(&call.name) {
return Ok(None);
}
let Some(value) = call.arguments.get("block_until_ms") else {
return Ok(None);
};
let integer = if let Some(value) = value.as_i64() {
value
} else {
let value = value.as_f64().ok_or_else(|| {
Error::Protocol(format!("{} block_until_ms must be an integer", call.name))
})?;
if !value.is_finite() || value.fract() != 0.0 {
return Err(Error::Protocol(format!(
"{} block_until_ms must be an integer",
call.name
)));
}
if value < i64::MIN as f64 || value > i64::MAX as f64 {
return Err(Error::Protocol(format!(
"{} block_until_ms is out of range",
call.name
)));
}
value as i64
};
if integer < 0 {
return Err(Error::Protocol(format!(
"{} block_until_ms is out of range",
call.name
)));
}
if value.as_i64().is_some() {
return Ok(None);
}
let mut normalized_call = call.clone();
normalized_call
.arguments
.as_object_mut()
.ok_or_else(|| Error::Protocol(format!("{} arguments must be a JSON object", call.name)))?
.insert("block_until_ms".into(), serde_json::Value::from(integer));
Ok(Some(normalized_call))
}
fn is_mcp_auth(call: &ToolCall) -> bool {
normalized(&call.name) == "callmcptool"
&& call
.arguments
.get("toolName")
.and_then(serde_json::Value::as_str)
.is_some_and(|tool| normalized(tool) == "mcpauth")
}
pub(super) async fn resume_interaction(
results: &ToolResultSender,
search: &WebSearch,
fetch: &WebFetch,
pending: PendingInteraction,
response: &pb::InteractionResponse,
) -> Result<InteractionContinuation> {
interaction::resume(results, search, fetch, pending, response).await
}
fn is_shell_tool(name: &str) -> bool {
matches!(normalized(name).as_str(), "shell" | "bash")
}
pub(super) fn normalized(name: &str) -> String {
name.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
@@ -0,0 +1,38 @@
//! Dispatches search Tool calls.
//! Cursor tool orchestration for application-owned Semble search.
use crate::{
cursor::tools::{
runtime::now_ms,
tool_call_result::{self as result, ToolResultSender},
},
model::ToolCall,
search,
store::Store,
Result,
};
use super::ToolStart;
pub(super) fn start(
results: &ToolResultSender,
call: &ToolCall,
store: Option<Store>,
) -> Result<ToolStart> {
let tool_name = super::normalized(&call.name);
let arguments = call.arguments.clone();
let call = call.clone();
let results = results.clone();
let started_at_ms = now_ms();
tokio::spawn(async move {
let output = search::execute_semble(&tool_name, arguments, store).await;
match result::semble(&call, started_at_ms, output) {
Ok(completion) => results.send(completion),
Err(error) => results.send_error(error),
}
});
Ok(ToolStart {
messages: Vec::new(),
completion: None,
})
}
@@ -0,0 +1,166 @@
//! Coordinates command execution Tool results.
mod output;
mod render;
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec as interaction},
model::ToolResult,
Error, Result,
};
use super::{gate, mcp_state, ReadImage, ToolCompletion};
use crate::cursor::tools::{
edit,
runtime::{ExecStage, PendingExec},
};
pub(crate) fn from_exec(
pending: PendingExec,
wire_result: &pb::exec_client_message::Message,
) -> Result<ToolCompletion> {
use pb::{exec_client_message::Message, tool_call::Tool};
let mut gated_shell = matches!(
wire_result,
Message::ShellResult(_) | Message::MiniSweAgentBashResult(_)
)
.then(|| wire_result.clone());
if let Some(message) = gated_shell.as_mut() {
gate::exec_message(message);
}
let wire_result = gated_shell.as_ref().unwrap_or(wire_result);
if let Message::McpStateExecResult(result) = wire_result {
return mcp_state::complete(pending, result);
}
let call = &pending.call;
let read_image = read_image(wire_result);
let (mut content, is_error) = output::output(wire_result, call)?;
if let Some(image) = &read_image {
content = format!("Read image file: {}", image.path);
}
let mut rendered = match &pending.stage {
ExecStage::DynamicMcp(definition) => {
interaction::render_dynamic_mcp(call, definition, false)
}
_ => interaction::render_tool_call(call, false)?,
};
match (rendered.tool.as_mut(), wire_result) {
(Some(Tool::ShellToolCall(tool)), Message::ShellResult(result))
| (Some(Tool::ShellToolCall(tool)), Message::MiniSweAgentBashResult(result)) => {
tool.result = Some(result.clone());
}
(Some(Tool::DeleteToolCall(tool)), Message::DeleteResult(result)) => {
tool.result = Some(result.clone());
}
(Some(Tool::GrepToolCall(tool)), Message::GrepResult(result)) => {
tool.result = Some(result.clone());
}
(Some(Tool::GlobToolCall(tool)), Message::GrepResult(result)) => {
tool.result = Some(render::glob(result)?);
}
(Some(Tool::ReadToolCall(tool)), Message::ReadResult(result))
| (Some(Tool::ReadToolCall(tool)), Message::RedactedReadResult(result)) => {
tool.result = Some(render::read(result, call)?);
}
(Some(Tool::ReadLintsToolCall(tool)), Message::DiagnosticsResult(result)) => {
tool.result = Some(render::diagnostics(result)?);
}
(Some(Tool::McpToolCall(tool)), Message::McpResult(result)) => {
tool.result = Some(render::mcp(result)?);
}
(Some(Tool::ReadMcpResourceToolCall(tool)), Message::ReadMcpResourceExecResult(result)) => {
tool.result = Some(result.clone());
}
(Some(Tool::TaskToolCall(tool)), Message::SubagentResult(result)) => {
tool.result = Some(render::task(result, call, pending.started_at_ms)?);
}
(Some(Tool::EditToolCall(tool)), Message::WriteResult(result)) => {
tool.result = Some(match (&pending.stage, result.result.as_ref()) {
(ExecStage::EditWrite(write), Some(pb::write_result::Result::Success(success))) => {
edit::success(success.path.clone(), write)
}
_ => render::write(result)?,
});
}
_ => {
return Err(Error::Protocol(format!(
"unexpected Exec result for tool {}",
call.name
)));
}
}
let tool = rendered.tool.ok_or_else(|| {
Error::Protocol(format!("tool {} has no Cursor representation", call.name))
})?;
Ok(ToolCompletion::new(
call,
pending.started_at_ms,
ToolResult {
call_id: call.call_id.clone(),
content,
is_error,
image: None,
},
tool,
)
.with_read_image(read_image))
}
fn read_image(message: &pb::exec_client_message::Message) -> Option<ReadImage> {
use pb::{exec_client_message::Message, read_result::Result, read_success::Output};
let result = match message {
Message::ReadResult(result) | Message::RedactedReadResult(result) => result,
_ => return None,
};
let Result::Success(success) = result.result.as_ref()? else {
return None;
};
let Output::Data(data) = success.output.as_ref()? else {
return None;
};
Some(ReadImage {
mime_type: image_mime_type(data)?.into(),
data: data.clone(),
path: success.path.clone(),
})
}
fn image_mime_type(data: &[u8]) -> Option<&'static str> {
let reader = image::ImageReader::new(std::io::Cursor::new(data))
.with_guessed_format()
.ok()?;
let format = reader.format()?;
let (width, height) = reader.into_dimensions().ok()?;
if width == 0 || height == 0 {
return None;
}
match format {
image::ImageFormat::Png => Some("image/png"),
image::ImageFormat::Jpeg => Some("image/jpeg"),
image::ImageFormat::Gif => Some("image/gif"),
image::ImageFormat::WebP => Some("image/webp"),
_ => None,
}
}
pub(crate) fn edit_failure(pending: PendingExec, error: String) -> Result<ToolCompletion> {
let call = &pending.call;
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::EditToolCall(mut tool)) = rendered.tool.take() else {
return Err(Error::Protocol(format!(
"{} is not an edit tool",
call.name
)));
};
tool.result = Some(edit::failure(edit::path(call)?, error.clone()));
Ok(ToolCompletion::new(
call,
pending.started_at_ms,
ToolResult {
call_id: call.call_id.clone(),
content: error,
is_error: true,
image: None,
},
pb::tool_call::Tool::EditToolCall(tool),
))
}
@@ -0,0 +1,415 @@
//! Parses and persists command execution output.
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolCall, Error, Result};
pub(super) fn output(
message: &pb::exec_client_message::Message,
call: &ToolCall,
) -> Result<(String, bool)> {
use pb::exec_client_message::Message;
match message {
Message::ShellResult(value) | Message::MiniSweAgentBashResult(value) => shell(value),
Message::ReadResult(value) | Message::RedactedReadResult(value) => read(value),
Message::WriteResult(value) => write(value),
Message::DeleteResult(value) => delete(value),
Message::GrepResult(value) => grep(value),
Message::DiagnosticsResult(value) => diagnostics(value),
Message::McpResult(value) => mcp(value),
Message::ReadMcpResourceExecResult(value) => read_mcp(value),
Message::SubagentResult(value) => task(value, call),
_ => Err(Error::Protocol(
"unsupported terminal ExecClientMessage".into(),
)),
}
}
fn shell(value: &pb::ShellResult) -> Result<(String, bool)> {
use pb::shell_result::Result as R;
let output = match value.result.as_ref().ok_or_else(|| missing("shell"))? {
R::Success(success) if value.is_background == Some(true) => {
let mut fields = vec![format!("shell_id={}", success.shell_id.unwrap_or_default())];
if let Some(pid) = success.pid.or(value.pid) {
fields.push(format!("pid={pid}"));
}
if let Some(folder) = value.terminals_folder.as_deref().filter(|v| !v.is_empty()) {
fields.push(format!("terminals_folder={folder}"));
}
let output = streams(&success.stdout, &success.stderr);
let prefix = format!("shell running in background {}", fields.join(" "));
return Ok((
if output == "shell completed without output" {
prefix
} else {
format!("{prefix}\n{output}")
},
false,
));
}
R::Success(success) => return Ok((streams(&success.stdout, &success.stderr), false)),
R::Failure(failure) => streams(&failure.stdout, &failure.stderr),
R::Timeout(timeout) => format!(
"shell timed out after {}ms in {}",
timeout.timeout_ms, timeout.working_directory
),
R::Rejected(rejected) => rejected.reason.clone(),
R::SpawnError(error) => error.error.clone(),
R::PermissionDenied(denied) => denied.error.clone(),
};
Ok((output, true))
}
fn streams(stdout: &str, stderr: &str) -> String {
match (stdout.is_empty(), stderr.is_empty()) {
(false, false) => format!("{stdout}\n\n<stderr>\n{stderr}\n</stderr>"),
(false, true) => stdout.into(),
(true, false) => stderr.into(),
(true, true) => "shell completed without output".into(),
}
}
fn read(value: &pb::ReadResult) -> Result<(String, bool)> {
use pb::{read_result::Result as R, read_success::Output};
match value.result.as_ref().ok_or_else(|| missing("read"))? {
R::Success(success) => Ok((
match success.output.as_ref() {
Some(Output::Content(text)) => text.clone(),
Some(Output::Data(bytes)) => format!("read binary bytes={}", bytes.len()),
None => format!("read success path={}", success.path),
},
false,
)),
R::Error(error) => Ok((error.error.clone(), true)),
R::Rejected(rejected) => Ok((rejected.reason.clone(), true)),
R::FileNotFound(value) => Ok((format!("file not found: {}", value.path), true)),
R::PermissionDenied(value) => Ok((format!("permission denied: {}", value.path), true)),
R::InvalidFile(value) => Ok((value.reason.clone(), true)),
}
}
fn write(value: &pb::WriteResult) -> Result<(String, bool)> {
use pb::write_result::Result as R;
match value.result.as_ref().ok_or_else(|| missing("write"))? {
R::Success(success) => Ok((
success.file_content_after_write.clone().unwrap_or_else(|| {
format!(
"write success path={} lines={}",
success.path, success.lines_created
)
}),
false,
)),
R::PermissionDenied(value) => Ok((value.error.clone(), true)),
R::NoSpace(value) => Ok((format!("no space left: {}", value.path), true)),
R::Error(value) => Ok((value.error.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
}
}
fn delete(value: &pb::DeleteResult) -> Result<(String, bool)> {
use pb::delete_result::Result as R;
match value.result.as_ref().ok_or_else(|| missing("delete"))? {
R::Success(value) => Ok((format!("delete success path={}", value.path), false)),
R::FileNotFound(value) => Ok((format!("file not found: {}", value.path), true)),
R::NotFile(value) => Ok((format!("not file: {}", value.path), true)),
R::PermissionDenied(value) => Ok((value.client_visible_error.clone(), true)),
R::FileBusy(value) => Ok((format!("file busy: {}", value.path), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::Error(value) => Ok((value.error.clone(), true)),
}
}
fn grep(value: &pb::GrepResult) -> Result<(String, bool)> {
use pb::grep_result::Result as R;
match value.result.as_ref().ok_or_else(|| missing("grep"))? {
R::Success(value) => Ok((grep_success(value), false)),
R::Error(value) => Ok((value.error.clone(), true)),
}
}
fn grep_success(value: &pb::GrepSuccess) -> String {
let mut lines = Vec::new();
if let Some(result) = &value.active_editor_result {
grep_union(result, &mut lines);
}
let mut workspaces = value.workspace_results.iter().collect::<Vec<_>>();
workspaces.sort_unstable_by_key(|(name, _)| *name);
for (_, result) in workspaces {
grep_union(result, &mut lines);
}
if lines.is_empty() {
format!(
"No matches found for pattern `{}` in {}",
value.pattern, value.path
)
} else {
lines.join("\n")
}
}
fn grep_union(value: &pb::GrepUnionResult, lines: &mut Vec<String>) {
use pb::grep_union_result::Result as R;
match value.result.as_ref() {
Some(R::Files(value)) => {
lines.extend(value.files.iter().cloned());
grep_truncation(
value.client_truncated,
value.ripgrep_truncated,
value.total_files,
"files",
lines,
);
}
Some(R::Count(value)) => {
lines.extend(
value
.counts
.iter()
.map(|count| format!("{}:{}", count.file, count.count)),
);
grep_truncation(
value.client_truncated,
value.ripgrep_truncated,
value.total_matches,
"matches",
lines,
);
}
Some(R::Content(value)) => {
for file in &value.matches {
lines.extend(file.matches.iter().map(|matched| {
let separator = if matched.is_context_line { '-' } else { ':' };
let truncated = if matched.content_truncated {
" [line truncated]"
} else {
""
};
format!(
"{}{separator}{}{separator}{}{truncated}",
file.file, matched.line_number, matched.content
)
}));
}
grep_truncation(
value.client_truncated,
value.ripgrep_truncated,
value.total_matched_lines,
"matched lines",
lines,
);
}
None => {}
}
}
fn grep_truncation(
client_truncated: bool,
ripgrep_truncated: bool,
total: i32,
unit: &str,
lines: &mut Vec<String>,
) {
if client_truncated || ripgrep_truncated {
lines.push(format!("[Results truncated; {total} total {unit}]"));
}
}
fn diagnostics(value: &pb::DiagnosticsResult) -> Result<(String, bool)> {
use pb::diagnostics_result::Result as R;
match value
.result
.as_ref()
.ok_or_else(|| missing("diagnostics"))?
{
R::Success(value) => Ok((diagnostics_success(value), false)),
R::Error(value) => Ok((value.error.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::FileNotFound(value) => Ok((format!("file not found: {}", value.path), true)),
R::PermissionDenied(value) => Ok((format!("permission denied: {}", value.path), true)),
}
}
fn diagnostics_success(value: &pb::DiagnosticsSuccess) -> String {
if value.diagnostics.is_empty() {
return format!("No diagnostics found in {}", value.path);
}
let mut lines = value
.diagnostics
.iter()
.map(|diagnostic| {
let location = diagnostic_location(&value.path, diagnostic.range.as_ref());
let mut labels = vec![diagnostic_severity(diagnostic.severity)];
if !diagnostic.source.is_empty() {
labels.push(diagnostic.source.as_str());
}
if !diagnostic.code.is_empty() {
labels.push(diagnostic.code.as_str());
}
if diagnostic.is_stale {
labels.push("stale");
}
format!(
"{}: [{}] {}",
location,
labels.join(" "),
diagnostic.message
)
})
.collect::<Vec<_>>();
if value.total_diagnostics != value.diagnostics.len() as i32 {
lines.push(format!(
"[Reported {} diagnostics; received {} details]",
value.total_diagnostics,
value.diagnostics.len()
));
}
lines.join("\n")
}
fn diagnostic_location(path: &str, range: Option<&pb::Range>) -> String {
let Some(range) = range else {
return path.into();
};
let Some(start) = &range.start else {
return path.into();
};
let mut location = format!(
"{}:{}:{}",
path,
start.line.saturating_add(1),
start.column.saturating_add(1)
);
if let Some(end) = &range.end {
location.push_str(&format!(
"-{}:{}",
end.line.saturating_add(1),
end.column.saturating_add(1)
));
}
location
}
fn diagnostic_severity(value: i32) -> &'static str {
match pb::DiagnosticSeverity::try_from(value) {
Ok(pb::DiagnosticSeverity::Error) => "error",
Ok(pb::DiagnosticSeverity::Warning) => "warning",
Ok(pb::DiagnosticSeverity::Information) => "information",
Ok(pb::DiagnosticSeverity::Hint) => "hint",
Ok(pb::DiagnosticSeverity::Unspecified) | Err(_) => "diagnostic",
}
}
fn mcp(value: &pb::McpResult) -> Result<(String, bool)> {
use pb::mcp_result::Result as R;
match value.result.as_ref().ok_or_else(|| missing("mcp"))? {
R::Success(value) => Ok((mcp_content(value)?, value.is_error)),
R::Error(value) => Ok((value.error.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::PermissionDenied(value) => Ok((value.error.clone(), true)),
R::ToolNotFound(value) => Ok((format!("MCP tool not found: {}", value.name), true)),
R::ServerNotFound(value) => Ok((format!("MCP server not found: {}", value.name), true)),
R::Approved(_) => Err(Error::Protocol("MCP approval is not terminal".into())),
}
}
fn mcp_content(success: &pb::McpSuccess) -> Result<String> {
let mut content = Vec::new();
for item in &success.content {
match item.content.as_ref() {
Some(pb::mcp_tool_result_content_item::Content::Text(text)) => {
if !text.text.is_empty() {
content.push(text.text.clone());
}
if let Some(location) = &text.output_location {
content.push(format!(
"MCP output file: {} ({} bytes, {} lines)",
location.file_path, location.size_bytes, location.line_count
));
}
}
Some(pb::mcp_tool_result_content_item::Content::Image(image)) => content.push(format!(
"MCP image: {} ({} bytes)",
image.mime_type,
image.data.len()
)),
None => {}
}
}
if let Some(structured) = &success.structured_content {
let value = serde_json::Value::Object(
structured
.fields
.iter()
.map(|(key, value)| (key.clone(), super::super::prost_json(value)))
.collect(),
);
content.push(serde_json::to_string_pretty(&value)?);
}
Ok(if content.is_empty() {
"MCP tool completed without content".into()
} else {
content.join("\n\n")
})
}
fn read_mcp(value: &pb::ReadMcpResourceExecResult) -> Result<(String, bool)> {
use pb::read_mcp_resource_exec_result::Result as R;
match value
.result
.as_ref()
.ok_or_else(|| missing("read MCP resource"))?
{
R::Success(value) => Ok((
match value.content.as_ref() {
Some(pb::read_mcp_resource_success::Content::Text(text)) => text.clone(),
Some(pb::read_mcp_resource_success::Content::Blob(blob)) => {
format!("read MCP resource blob={}", blob.len())
}
None => format!("read MCP resource uri={}", value.uri),
},
false,
)),
R::Error(value) => Ok((value.error.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::NotFound(value) => Ok((format!("MCP resource not found: {}", value.uri), true)),
}
}
fn task(value: &pb::SubagentResult, call: &ToolCall) -> Result<(String, bool)> {
use pb::subagent_result::Result as R;
match value.result.as_ref().ok_or_else(|| missing("subagent"))? {
R::Success(value) if creates_subagent(call) => {
let name = call
.arguments
.get("description")
.and_then(serde_json::Value::as_str)
.filter(|name| !name.is_empty())
.ok_or_else(|| Error::Protocol("Task call is missing description".into()))?;
if value.agent_id.is_empty() {
return Err(Error::Protocol("Task result is missing agent_id".into()));
}
let identity = format!("Subagent name: {name}\nSubagent ID: {}", value.agent_id);
let content = value
.final_message
.as_deref()
.filter(|message| !message.is_empty())
.map_or(identity.clone(), |message| {
format!("{identity}\n\n{message}")
});
Ok((content, false))
}
R::Success(value) => Ok((value.final_message.clone().unwrap_or_default(), false)),
R::Error(value) => Ok((value.error.clone(), true)),
}
}
fn creates_subagent(call: &ToolCall) -> bool {
matches!(
call.arguments
.get("resume")
.and_then(serde_json::Value::as_str),
None | Some("self")
)
}
fn missing(name: &str) -> Error {
Error::Protocol(format!("{name} returned no result"))
}
@@ -0,0 +1,255 @@
//! Renders command execution output for Cursor.
use serde_json::Value;
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolCall, Error, Result};
pub(super) fn read(result: &pb::ReadResult, call: &ToolCall) -> Result<pb::ReadToolResult> {
use pb::{read_result::Result as Input, read_tool_result::Result as Output};
let result = match result.result.as_ref() {
Some(Input::Success(success)) => Output::Success(pb::ReadToolSuccess {
is_empty: match success.output.as_ref() {
Some(pb::read_success::Output::Content(content)) => content.is_empty(),
Some(pb::read_success::Output::Data(data)) => data.is_empty(),
None => true,
},
exceeded_limit: success.truncated,
total_lines: success.total_lines.max(0) as u32,
file_size: success.file_size.max(0).min(u32::MAX as i64) as u32,
path: success.path.clone(),
read_range: read_range(call),
include_line_numbers: call
.arguments
.get("include_line_numbers")
.and_then(Value::as_bool),
output: success.output.as_ref().map(|output| match output {
pb::read_success::Output::Content(content) => {
pb::read_tool_success::Output::Content(content.clone())
}
pb::read_success::Output::Data(data) => {
pb::read_tool_success::Output::Data(data.clone())
}
}),
..Default::default()
}),
Some(Input::Error(value)) => error_read(&value.error),
Some(Input::Rejected(value)) => error_read(&value.reason),
Some(Input::FileNotFound(value)) => error_read(&format!("file not found: {}", value.path)),
Some(Input::PermissionDenied(value)) => {
error_read(&format!("permission denied: {}", value.path))
}
Some(Input::InvalidFile(value)) => error_read(&value.reason),
None => return Err(missing("read")),
};
Ok(pb::ReadToolResult {
result: Some(result),
})
}
fn error_read(message: &str) -> pb::read_tool_result::Result {
pb::read_tool_result::Result::Error(pb::ReadToolError {
error_message: message.into(),
})
}
fn read_range(call: &ToolCall) -> Option<pb::ReadRange> {
let start_line = call
.arguments
.get("offset")
.and_then(Value::as_u64)
.unwrap_or(0) as u32;
let limit = call
.arguments
.get("limit")
.and_then(Value::as_u64)
.map(|value| value as u32)?;
Some(pb::ReadRange {
start_line,
end_line: start_line.saturating_add(limit),
})
}
pub(super) fn write(result: &pb::WriteResult) -> Result<pb::EditResult> {
use pb::{edit_result::Result as Output, write_result::Result as Input};
let result = match result.result.as_ref() {
Some(Input::Success(success)) => Output::Success(pb::EditSuccess {
path: success.path.clone(),
after_full_file_content: success.file_content_after_write.clone().unwrap_or_default(),
..Default::default()
}),
Some(Input::PermissionDenied(value)) => {
Output::WritePermissionDenied(pb::EditWritePermissionDenied {
path: value.path.clone(),
error: value.error.clone(),
is_readonly: value.is_readonly,
})
}
Some(Input::NoSpace(value)) => edit_error(&value.path, "no space left"),
Some(Input::Error(value)) => edit_error(&value.path, &value.error),
Some(Input::Rejected(value)) => Output::Rejected(pb::EditRejected {
path: value.path.clone(),
reason: value.reason.clone(),
}),
None => return Err(missing("write")),
};
Ok(pb::EditResult {
result: Some(result),
})
}
fn edit_error(path: &str, message: &str) -> pb::edit_result::Result {
pb::edit_result::Result::Error(pb::EditError {
path: path.into(),
error: message.into(),
model_visible_error: Some(message.into()),
})
}
pub(super) fn diagnostics(result: &pb::DiagnosticsResult) -> Result<pb::ReadLintsToolResult> {
use pb::{diagnostics_result::Result as Input, read_lints_tool_result::Result as Output};
let result = match result.result.as_ref() {
Some(Input::Success(success)) => {
let diagnostics = success
.diagnostics
.iter()
.map(|diagnostic| pb::DiagnosticItem {
severity: diagnostic.severity,
range: diagnostic.range.as_ref().map(|range| pb::DiagnosticRange {
start: range.start,
end: range.end,
}),
message: diagnostic.message.clone(),
source: diagnostic.source.clone(),
code: diagnostic.code.clone(),
is_stale: diagnostic.is_stale,
})
.collect::<Vec<_>>();
Output::Success(pb::ReadLintsToolSuccess {
file_diagnostics: vec![pb::FileDiagnostics {
path: success.path.clone(),
diagnostics_count: diagnostics.len() as i32,
diagnostics,
}],
total_files: 1,
total_diagnostics: success.total_diagnostics,
})
}
Some(Input::Error(value)) => lint_error(&value.error),
Some(Input::Rejected(value)) => lint_error(&value.reason),
Some(Input::FileNotFound(value)) => lint_error(&format!("file not found: {}", value.path)),
Some(Input::PermissionDenied(value)) => {
lint_error(&format!("permission denied: {}", value.path))
}
None => return Err(missing("diagnostics")),
};
Ok(pb::ReadLintsToolResult {
result: Some(result),
})
}
fn lint_error(message: &str) -> pb::read_lints_tool_result::Result {
pb::read_lints_tool_result::Result::Error(pb::ReadLintsToolError {
error_message: message.into(),
})
}
pub(super) fn mcp(result: &pb::McpResult) -> Result<pb::McpToolResult> {
use pb::{mcp_result::Result as Input, mcp_tool_result::Result as Output};
let result = match result.result.as_ref() {
Some(Input::Success(value)) => Output::Success(value.clone()),
Some(Input::Error(value)) => mcp_error(&value.error),
Some(Input::Rejected(value)) => Output::Rejected(value.clone()),
Some(Input::PermissionDenied(value)) => Output::PermissionDenied(value.clone()),
Some(Input::ToolNotFound(value)) => {
mcp_error(&format!("MCP tool not found: {}", value.name))
}
Some(Input::ServerNotFound(value)) => {
mcp_error(&format!("MCP server not found: {}", value.name))
}
Some(Input::Approved(_)) => {
return Err(Error::Protocol("MCP approval is not terminal".into()))
}
None => return Err(missing("MCP")),
};
Ok(pb::McpToolResult {
result: Some(result),
})
}
fn mcp_error(message: &str) -> pb::mcp_tool_result::Result {
pb::mcp_tool_result::Result::Error(pb::McpToolError {
error: message.into(),
read_tool_def_reminder: String::new(),
})
}
pub(super) fn task(
result: &pb::SubagentResult,
call: &crate::model::ToolCall,
started_at_ms: u64,
) -> Result<pb::TaskResult> {
use pb::{subagent_result::Result as Input, task_result::Result as Output};
let result = match result.result.as_ref() {
Some(Input::Success(value)) => {
let is_background = value.background_reason
!= pb::SubagentBackgroundReason::Unspecified as i32
|| call
.arguments
.get("run_in_background")
.and_then(serde_json::Value::as_bool)
== Some(true);
Output::Success(pb::TaskSuccess {
agent_id: Some(value.agent_id.clone()),
is_background,
duration_ms: Some(
crate::cursor::tools::runtime::now_ms().saturating_sub(started_at_ms),
),
result_suffix: value.final_message.clone(),
background_reason: value.background_reason,
transcript_path: value.transcript_path.clone(),
..Default::default()
})
}
Some(Input::Error(value)) => Output::Error(pb::TaskError {
error: value.error.clone(),
}),
None => return Err(missing("subagent")),
};
Ok(pb::TaskResult {
result: Some(result),
})
}
pub(super) fn glob(result: &pb::GrepResult) -> Result<pb::GlobToolResult> {
use pb::{glob_tool_result::Result as Output, grep_result::Result as Input};
let result = match result.result.as_ref() {
Some(Input::Success(success)) => {
let files = success
.active_editor_result
.iter()
.chain(success.workspace_results.values())
.find_map(|result| match result.result.as_ref() {
Some(pb::grep_union_result::Result::Files(files)) => Some(files),
_ => None,
});
Output::Success(pb::GlobToolSuccess {
pattern: success.pattern.clone(),
path: success.path.clone(),
files: files.map(|value| value.files.clone()).unwrap_or_default(),
total_files: files.map_or(0, |value| value.total_files),
client_truncated: files.is_some_and(|value| value.client_truncated),
ripgrep_truncated: files.is_some_and(|value| value.ripgrep_truncated),
})
}
Some(Input::Error(value)) => Output::Error(pb::GlobToolError {
error: value.error.clone(),
}),
None => return Err(missing("glob")),
};
Ok(pb::GlobToolResult {
result: Some(result),
})
}
fn missing(name: &str) -> Error {
Error::Protocol(format!("{name} returned no result"))
}
@@ -0,0 +1,687 @@
//! Correlates Tool completion events and gates final result delivery.
use std::collections::BTreeMap;
use crate::{cursor::protocol::proto::agent::v1 as pb, model::limit_tool_result_text};
const KIB: usize = 1024;
const READ_CONTENT_LIMIT: usize = 64 * KIB;
const READ_BINARY_LIMIT: usize = 32 * KIB;
const SHELL_STREAM_LIMIT: usize = 16 * KIB;
const SHELL_INTERLEAVED_LIMIT: usize = 32 * KIB;
const GREP_CONTENT_LIMIT: usize = 32 * KIB;
const GREP_MATCH_LIMIT: usize = 2 * KIB;
const GREP_MATCHES_PER_FILE: usize = 100;
const GREP_TOTAL_MATCHES: usize = 300;
const GREP_LIST_LIMIT: usize = 300;
const GLOB_FILE_LIMIT: usize = 200;
const EDIT_RESULT_LIMIT: usize = 32 * KIB;
const PATCH_EDIT_RESULT_LIMIT: usize = 4 * KIB;
const MCP_TEXT_LIMIT: usize = 32 * KIB;
const MCP_CONTENT_ITEM_LIMIT: usize = 20;
const MCP_STRUCTURED_LIMIT: usize = 32 * KIB;
const MCP_BINARY_LIMIT: usize = 32 * KIB;
const MCP_RESOURCE_LIMIT: usize = 200;
const MCP_RESOURCE_DESCRIPTION_LIMIT: usize = KIB;
const WEB_FETCH_LIMIT: usize = 32 * KIB;
const WEB_SEARCH_LIMIT: usize = 16 * KIB;
const WEB_SEARCH_TITLE_LIMIT: usize = 512;
const WEB_SEARCH_SNIPPET_LIMIT: usize = 2 * KIB;
pub(super) fn tool_completion(
tool_name: &str,
tool: &mut pb::tool_call::Tool,
content: &mut String,
) {
use pb::tool_call::Tool;
match tool {
Tool::ShellToolCall(tool) => gate_shell(tool),
Tool::GrepToolCall(tool) => gate_grep(tool),
Tool::GlobToolCall(tool) => gate_glob(tool),
Tool::ReadToolCall(tool) => gate_read(tool),
Tool::EditToolCall(tool) => gate_edit(tool_name, tool),
Tool::McpToolCall(tool) => gate_mcp(tool),
Tool::ListMcpResourcesToolCall(tool) => gate_mcp_resources(tool),
Tool::ReadMcpResourceToolCall(tool) => gate_mcp_resource(tool),
Tool::GetMcpToolsToolCall(tool) => gate_mcp_tools(tool),
Tool::WebFetchToolCall(tool) => gate_web_fetch(tool),
Tool::WebSearchToolCall(tool) => gate_web_search(tool),
Tool::GenerateImageToolCall(tool) => gate_generate_image(tool),
_ => {}
}
*content = limit_tool_result_text(tool_name, content);
}
pub(super) fn exec_message(message: &mut pb::exec_client_message::Message) {
use pb::exec_client_message::Message;
match message {
Message::ShellResult(result) | Message::MiniSweAgentBashResult(result) => {
gate_shell_result(result)
}
_ => {}
}
}
fn gate_shell(tool: &mut pb::ShellToolCall) {
if let Some(result) = tool.result.as_mut() {
gate_shell_result(result);
}
}
fn gate_shell_result(result: &mut pb::ShellResult) {
use pb::shell_result::Result;
match result.result.as_mut() {
Some(Result::Success(success)) => {
success.stdout = truncate_edges("Shell stdout", &success.stdout, SHELL_STREAM_LIMIT);
success.stderr = truncate_edges("Shell stderr", &success.stderr, SHELL_STREAM_LIMIT);
if let Some(interleaved) = success.interleaved_output.as_mut() {
*interleaved = truncate_edges(
"Shell interleaved output",
interleaved,
SHELL_INTERLEAVED_LIMIT,
);
}
}
Some(Result::Failure(failure)) => {
failure.stdout = truncate_edges("Shell stdout", &failure.stdout, SHELL_STREAM_LIMIT);
failure.stderr = truncate_edges("Shell stderr", &failure.stderr, SHELL_STREAM_LIMIT);
if let Some(interleaved) = failure.interleaved_output.as_mut() {
*interleaved = truncate_edges(
"Shell interleaved output",
interleaved,
SHELL_INTERLEAVED_LIMIT,
);
}
}
_ => {}
}
}
fn gate_read(tool: &mut pb::ReadToolCall) {
let Some(pb::read_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let Some(output) = success.output.as_mut() else {
return;
};
match output {
pb::read_tool_success::Output::Content(value) => {
let next = truncate_text("Read", value, READ_CONTENT_LIMIT);
if next != *value {
*value = next;
success.exceeded_limit = true;
}
}
pb::read_tool_success::Output::Data(value) if value.len() > READ_BINARY_LIMIT => {
let notice = truncation_notice("Read binary data", READ_BINARY_LIMIT, 0, value.len());
success.output = Some(pb::read_tool_success::Output::Content(notice));
success.exceeded_limit = true;
}
_ => {}
}
}
fn gate_glob(tool: &mut pb::GlobToolCall) {
let Some(pb::glob_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let original = success.files.len();
if original <= GLOB_FILE_LIMIT {
if success.total_files <= 0 {
success.total_files = original as i32;
}
return;
}
success.files.truncate(GLOB_FILE_LIMIT);
success.total_files = success.total_files.max(original as i32);
success.client_truncated = true;
}
fn gate_grep(tool: &mut pb::GrepToolCall) {
let Some(pb::grep_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let mut budget = GrepBudget {
content_bytes: GREP_CONTENT_LIMIT,
matches: GREP_TOTAL_MATCHES,
};
let mut workspace_names = success
.workspace_results
.keys()
.cloned()
.collect::<Vec<_>>();
workspace_names.sort_unstable();
for name in workspace_names {
if let Some(result) = success.workspace_results.get_mut(&name) {
gate_grep_union(result, &mut budget);
}
}
if let Some(result) = success.active_editor_result.as_mut() {
gate_grep_union(result, &mut budget);
}
}
struct GrepBudget {
content_bytes: usize,
matches: usize,
}
fn gate_grep_union(result: &mut pb::GrepUnionResult, budget: &mut GrepBudget) {
use pb::grep_union_result::Result;
match result.result.as_mut() {
Some(Result::Content(content)) => gate_grep_content(content, budget),
Some(Result::Files(files)) => {
let original = files.files.len();
if original > GREP_LIST_LIMIT {
files.files.truncate(GREP_LIST_LIMIT);
files.client_truncated = true;
}
if files.total_files <= 0 {
files.total_files = original as i32;
}
}
Some(Result::Count(counts)) => {
let original = counts.counts.len();
if original > GREP_LIST_LIMIT {
counts.counts.truncate(GREP_LIST_LIMIT);
counts.client_truncated = true;
}
if counts.total_files <= 0 {
counts.total_files = original as i32;
}
}
None => {}
}
}
fn gate_grep_content(content: &mut pb::GrepContentResult, budget: &mut GrepBudget) {
if content
.matches
.iter()
.flat_map(|file| &file.matches)
.any(is_grep_notice)
{
return;
}
let original_bytes = grep_content_bytes(&content.matches);
let original_files = content.matches.len();
let mut truncated = false;
let mut files = Vec::with_capacity(original_files);
for file in &content.matches {
if budget.matches == 0 || budget.content_bytes == 0 {
truncated = true;
break;
}
let mut next = pb::GrepFileMatch {
file: file.file.clone(),
matches: Vec::new(),
};
for matched in &file.matches {
if is_grep_notice(matched) {
next.matches.push(matched.clone());
continue;
}
if next.matches.len() >= GREP_MATCHES_PER_FILE
|| budget.matches == 0
|| budget.content_bytes == 0
{
truncated = true;
break;
}
let mut next_match = matched.clone();
let original = next_match.content.clone();
next_match.content = truncate_text("Grep match", &original, GREP_MATCH_LIMIT);
if next_match.content != original {
next_match.content_truncated = true;
truncated = true;
}
if next_match.content.len() > budget.content_bytes {
next_match.content =
truncate_text("Grep", &next_match.content, budget.content_bytes);
next_match.content_truncated = true;
truncated = true;
}
if next_match.content.trim().is_empty() {
truncated = true;
break;
}
budget.content_bytes -= next_match.content.len();
budget.matches -= 1;
next.matches.push(next_match);
}
if next.matches.len() < file.matches.len() {
truncated = true;
}
if !next.matches.is_empty() {
files.push(next);
}
}
if files.len() < original_files {
truncated = true;
}
if truncated {
content.client_truncated = true;
add_grep_notice(&mut files, original_bytes);
}
content.matches = files;
}
fn add_grep_notice(files: &mut Vec<pb::GrepFileMatch>, original_bytes: usize) {
if files
.iter()
.flat_map(|file| &file.matches)
.any(is_grep_notice)
{
return;
}
loop {
let used = grep_content_bytes(files);
let notice = truncation_notice("Grep", GREP_CONTENT_LIMIT, used, original_bytes);
if used.saturating_add(notice.len()) <= GREP_CONTENT_LIMIT {
let matched = pb::GrepContentMatch {
line_number: 0,
content: notice,
content_truncated: true,
is_context_line: true,
};
if let Some(file) = files.last_mut() {
file.matches.push(matched);
} else {
files.push(pb::GrepFileMatch {
file: "[truncated]".into(),
matches: vec![matched],
});
}
return;
}
let Some(file) = files.last_mut() else {
return;
};
file.matches.pop();
if file.matches.is_empty() {
files.pop();
}
}
}
fn is_grep_notice(matched: &pb::GrepContentMatch) -> bool {
matched.line_number == 0
&& matched.content_truncated
&& matched
.content
.starts_with("[truncated: Grep result exceeded")
}
fn grep_content_bytes(files: &[pb::GrepFileMatch]) -> usize {
files
.iter()
.flat_map(|file| &file.matches)
.map(|matched| matched.content.len())
.sum()
}
fn gate_edit(tool_name: &str, tool: &mut pb::EditToolCall) {
let Some(pb::edit_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
let limit = match tool_name.trim() {
"PatchEdit" | "PatchEditLines" | "PatchEditSpan" | "StrReplace" => PATCH_EDIT_RESULT_LIMIT,
_ => EDIT_RESULT_LIMIT,
};
if let Some(diff) = success.diff_string.as_mut() {
*diff = truncate_text(tool_name, diff, limit);
success.before_full_file_content = None;
success.after_full_file_content.clear();
} else {
success.before_full_file_content = None;
success.after_full_file_content =
truncate_text(tool_name, &success.after_full_file_content, limit);
}
}
fn gate_mcp(tool: &mut pb::McpToolCall) {
let Some(pb::mcp_tool_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if success.content.iter().any(is_mcp_notice) {
return;
}
let mut notices = Vec::new();
if structured_json_len(&success.structured_content) > MCP_STRUCTURED_LIMIT {
let original = structured_json_len(&success.structured_content);
success.structured_content = truncated_struct(original, MCP_STRUCTURED_LIMIT);
notices.push(truncation_notice(
"MCP structured_content",
MCP_STRUCTURED_LIMIT,
0,
original,
));
}
let original_items = success.content.len();
if original_items > MCP_CONTENT_ITEM_LIMIT {
success.content.truncate(MCP_CONTENT_ITEM_LIMIT);
notices.push(format!(
"[truncated: MCP content items exceeded {MCP_CONTENT_ITEM_LIMIT} items; showing {MCP_CONTENT_ITEM_LIMIT} of {original_items} items]"
));
}
let mut remaining_text = MCP_TEXT_LIMIT;
let mut content = Vec::with_capacity(success.content.len() + notices.len());
for mut item in std::mem::take(&mut success.content) {
// MCP images are sent to the client as inline binary data. Truncating
// an encoded image at an arbitrary byte boundary corrupts the image
// and makes the client's image/screenshot fallback fail. The model
// receives only the textual MCP summary below, which is bounded by
// MCP_TEXT_LIMIT, so the image does not need this text-result gate.
if let Some(pb::mcp_tool_result_content_item::Content::Text(text)) = item.content.as_mut() {
let original = text.text.clone();
let next = truncate_text("MCP content item", &original, MCP_TEXT_LIMIT);
if remaining_text == 0 {
notices.push(truncation_notice(
"MCP text",
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT,
MCP_TEXT_LIMIT.saturating_add(original.len()),
));
continue;
}
text.text = truncate_text("MCP text", &next, remaining_text);
remaining_text = remaining_text.saturating_sub(text.text.len());
}
content.push(item);
}
content.extend(notices.into_iter().map(mcp_notice));
success.content = content;
}
fn mcp_notice(text: String) -> pb::McpToolResultContentItem {
pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text,
output_location: None,
},
)),
}
}
fn is_mcp_notice(item: &pb::McpToolResultContentItem) -> bool {
matches!(
item.content.as_ref(),
Some(pb::mcp_tool_result_content_item::Content::Text(text))
if text.text.starts_with("[truncated:")
)
}
fn structured_json_len(value: &Option<prost_types::Struct>) -> usize {
value
.as_ref()
.and_then(|value| {
serde_json::to_vec(&serde_json::Value::Object(
value
.fields
.iter()
.map(|(key, value)| (key.clone(), super::prost_json(value)))
.collect(),
))
.ok()
})
.map_or(0, |value| value.len())
}
fn truncated_struct(original: usize, limit: usize) -> Option<prost_types::Struct> {
Some(prost_types::Struct {
fields: BTreeMap::from([
("_truncated".into(), prost_bool(true)),
("original_json_bytes".into(), prost_number(original as f64)),
("limit_bytes".into(), prost_number(limit as f64)),
]),
})
}
fn prost_bool(value: bool) -> prost_types::Value {
prost_types::Value {
kind: Some(prost_types::value::Kind::BoolValue(value)),
}
}
fn prost_number(value: f64) -> prost_types::Value {
prost_types::Value {
kind: Some(prost_types::value::Kind::NumberValue(value)),
}
}
fn gate_mcp_resources(tool: &mut pb::ListMcpResourcesToolCall) {
let Some(pb::list_mcp_resources_exec_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if success
.resources
.iter()
.any(|resource| resource.uri == "truncated:list-mcp-resources")
{
return;
}
let original = success.resources.len();
success.resources.truncate(MCP_RESOURCE_LIMIT);
for resource in &mut success.resources {
if let Some(description) = resource.description.as_mut() {
*description = truncate_text(
"MCP resource description",
description,
MCP_RESOURCE_DESCRIPTION_LIMIT,
);
}
}
if success.resources.len() < original {
success
.resources
.push(pb::list_mcp_resources_exec_result::McpResource {
uri: "truncated:list-mcp-resources".into(),
name: Some("truncated".into()),
description: Some(truncation_notice(
"ListMcpResources",
MCP_TEXT_LIMIT,
success.resources.len(),
original,
)),
..Default::default()
});
}
}
fn gate_mcp_resource(tool: &mut pb::ReadMcpResourceToolCall) {
let Some(pb::read_mcp_resource_exec_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
match success.content.as_mut() {
Some(pb::read_mcp_resource_success::Content::Text(text)) => {
*text = truncate_text("FetchMcpResource", text, MCP_TEXT_LIMIT);
}
Some(pb::read_mcp_resource_success::Content::Blob(blob))
if blob.len() > MCP_BINARY_LIMIT =>
{
let notice =
truncation_notice("FetchMcpResource blob", MCP_BINARY_LIMIT, 0, blob.len());
success.content = Some(pb::read_mcp_resource_success::Content::Text(notice));
}
_ => {}
}
}
fn gate_mcp_tools(tool: &mut pb::GetMcpToolsToolCall) {
let Some(pb::get_mcp_tools_agent_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
success.content = truncate_text("GetMcpTools", &success.content, MCP_TEXT_LIMIT);
}
fn gate_web_fetch(tool: &mut pb::WebFetchToolCall) {
let Some(pb::web_fetch_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
success.markdown = truncate_text("WebFetch", &success.markdown, WEB_FETCH_LIMIT);
}
fn gate_web_search(tool: &mut pb::WebSearchToolCall) {
let Some(pb::web_search_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
for reference in &mut success.references {
reference.title =
truncate_text("WebSearch title", &reference.title, WEB_SEARCH_TITLE_LIMIT);
reference.chunk = truncate_text(
"WebSearch snippet",
&reference.chunk,
WEB_SEARCH_SNIPPET_LIMIT,
);
}
let original = web_search_bytes(&success.references);
while success.references.len() > 1 && web_search_bytes(&success.references) > WEB_SEARCH_LIMIT {
success.references.pop();
}
if original > WEB_SEARCH_LIMIT {
let total = web_search_bytes(&success.references);
if let Some(reference) = success.references.last_mut() {
let other = total.saturating_sub(reference.chunk.len());
let notice = truncation_notice(
"WebSearch",
WEB_SEARCH_LIMIT,
WEB_SEARCH_LIMIT.saturating_sub(other),
original,
);
let available = WEB_SEARCH_LIMIT.saturating_sub(other + notice.len() + 2);
reference.chunk = format!(
"{}\n\n{notice}",
utf8_prefix(&reference.chunk, available).trim_end_matches('\n')
);
}
}
}
fn web_search_bytes(references: &[pb::WebSearchReference]) -> usize {
references
.iter()
.map(|reference| reference.title.len() + reference.url.len() + reference.chunk.len())
.sum()
}
fn gate_generate_image(tool: &mut pb::GenerateImageToolCall) {
let Some(pb::generate_image_result::Result::Success(success)) = tool
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return;
};
if !success.image_data.trim().is_empty()
&& !success
.image_data
.starts_with("[base64 image data omitted from replay; bytes=")
{
let original = success.image_data.trim().len();
success.image_data = format!("[base64 image data omitted from replay; bytes={original}]");
}
}
fn truncate_text(tool_name: &str, content: &str, limit: usize) -> String {
if content.len() <= limit {
return content.to_string();
}
let original = content.len();
let mut shown = limit;
loop {
let notice = format!(
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; showing {shown} of {original} bytes]"
);
let available = limit.saturating_sub(notice.len());
let kept = utf8_prefix(content, available);
if kept.len() == shown {
return format!("{}{notice}", kept.trim_end_matches('\n'));
}
shown = kept.len();
}
}
fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
if content.len() <= limit {
return content.to_string();
}
let original = content.len();
let mut shown = limit;
loop {
let notice = format!(
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; omitted middle; showing {shown} of {original} bytes]\n\n"
);
let available = limit.saturating_sub(notice.len());
let head = utf8_prefix(content, available / 2);
let tail = utf8_suffix(content, available.saturating_sub(head.len()));
let next_shown = head.len().saturating_add(tail.len());
if next_shown == shown {
return format!("{head}{notice}{tail}");
}
shown = next_shown;
}
}
fn truncation_notice(tool_name: &str, limit: usize, shown: usize, original: usize) -> String {
format!(
"[truncated: {tool_name} result exceeded {limit} bytes; showing {shown} of {original} bytes]"
)
}
fn utf8_prefix(value: &str, limit: usize) -> &str {
let mut end = limit.min(value.len());
while end > 0 && !value.is_char_boundary(end) {
end -= 1;
}
&value[..end]
}
fn utf8_suffix(value: &str, limit: usize) -> &str {
let mut start = value.len().saturating_sub(limit);
while start < value.len() && !value.is_char_boundary(start) {
start += 1;
}
&value[start..]
}
@@ -0,0 +1,330 @@
//! Converts Cursor interaction completions into Tool results.
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec as interaction},
search::{FetchedPage, SearchHit},
Error, Result,
};
use super::ToolCompletion;
use crate::cursor::tools::runtime::PendingInteraction;
pub(crate) fn from_interaction(
pending: PendingInteraction,
response: &pb::InteractionResponse,
) -> Result<ToolCompletion> {
use pb::{interaction_response::Result as Response, tool_call::Tool};
let call = &pending.call;
let mut rendered = interaction::render_tool_call(call, false)?;
let (output, is_error) = match (rendered.tool.as_mut(), response.result.as_ref()) {
(
Some(Tool::AskQuestionToolCall(tool)),
Some(Response::AskQuestionInteractionResponse(value)),
) => {
let result = value
.result
.clone()
.ok_or_else(|| missing("ask question"))?;
let output = ask_output(&result)?;
tool.result = Some(result);
output
}
(
Some(Tool::CreatePlanToolCall(tool)),
Some(Response::CreatePlanRequestResponse(value)),
) => {
let result = value.result.clone().ok_or_else(|| missing("create plan"))?;
let output = create_plan_output(&result)?;
tool.result = Some(result);
output
}
(
Some(Tool::SwitchModeToolCall(tool)),
Some(Response::SwitchModeRequestResponse(value)),
) => {
let (result, output) = switch_mode_result(value)?;
tool.result = Some(result);
output
}
(Some(Tool::WebSearchToolCall(tool)), Some(Response::WebSearchRequestResponse(value))) => {
match value
.result
.as_ref()
.ok_or_else(|| missing("web search approval"))?
{
pb::web_search_request_response::Result::Rejected(rejected) => {
tool.result = Some(pb::WebSearchResult {
result: Some(pb::web_search_result::Result::Rejected(
pb::WebSearchRejected {
reason: rejected.reason.clone(),
},
)),
});
(rejected.reason.clone(), true)
}
pb::web_search_request_response::Result::Approved(_) => {
return Err(Error::Protocol(
"WebSearch approval reached terminal response decoding".into(),
));
}
}
}
(Some(Tool::WebFetchToolCall(tool)), Some(Response::WebFetchRequestResponse(value))) => {
match value
.result
.as_ref()
.ok_or_else(|| missing("web fetch approval"))?
{
pb::web_fetch_request_response::Result::Rejected(rejected) => {
tool.result = Some(pb::WebFetchResult {
result: Some(pb::web_fetch_result::Result::Rejected(
pb::WebFetchRejected {
reason: rejected.reason.clone(),
},
)),
});
(rejected.reason.clone(), true)
}
pb::web_fetch_request_response::Result::Approved(_) => {
return Err(Error::Protocol(
"WebFetch approval is not a terminal tool result".into(),
));
}
}
}
(
Some(Tool::GenerateImageToolCall(tool)),
Some(Response::GenerateImageRequestResponse(value)),
) => match value
.result
.as_ref()
.ok_or_else(|| missing("generate image approval"))?
{
pb::generate_image_request_response::Result::Rejected(rejected) => {
tool.result = Some(pb::GenerateImageResult {
result: Some(pb::generate_image_result::Result::Error(
pb::GenerateImageError {
error: rejected.reason.clone(),
},
)),
});
(rejected.reason.clone(), true)
}
pb::generate_image_request_response::Result::Approved(_) => {
return Err(Error::Provider(
"GenerateImage requires a configured server-side image executor".into(),
));
}
},
(Some(Tool::McpAuthToolCall(tool)), Some(Response::McpAuthRequestResponse(value))) => {
let server_identifier = tool
.args
.as_ref()
.map(|args| args.server_identifier.clone())
.unwrap_or_default();
let result = match value
.result
.as_ref()
.ok_or_else(|| missing("MCP authentication"))?
{
pb::mcp_auth_request_response::Result::Approved(_) => {
pb::mcp_auth_result::Result::Success(pb::McpAuthSuccess {
server_identifier: server_identifier.clone(),
})
}
pb::mcp_auth_request_response::Result::Rejected(rejected) => {
pb::mcp_auth_result::Result::Rejected(pb::McpAuthRejected {
reason: rejected.reason.clone(),
})
}
};
let (output, is_error) = match &result {
pb::mcp_auth_result::Result::Success(_) => (
format!("Authenticated MCP server {server_identifier}"),
false,
),
pb::mcp_auth_result::Result::Rejected(rejected) => (rejected.reason.clone(), true),
pb::mcp_auth_result::Result::Error(error) => (error.error.clone(), true),
};
tool.result = Some(pb::McpAuthResult {
result: Some(result),
});
(output, is_error)
}
_ => {
return Err(Error::Protocol(format!(
"unexpected InteractionResponse for tool {}",
call.name
)));
}
};
ToolCompletion::from_rendered(call, pending.started_at_ms, output, is_error, rendered)
}
pub(crate) fn complete_web_search(
pending: PendingInteraction,
outcome: std::result::Result<Vec<SearchHit>, String>,
) -> Result<ToolCompletion> {
let call = &pending.call;
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::WebSearchToolCall(tool)) = rendered.tool.as_mut() else {
return Err(Error::Protocol(format!(
"tool {} is not WebSearch",
call.name
)));
};
let (output, is_error) = match outcome {
Ok(hits) => {
let output = hits
.iter()
.enumerate()
.map(|(index, hit)| {
format!(
"{}. {}\nURL: {}\n{}",
index + 1,
hit.title,
hit.url,
hit.chunk
)
})
.collect::<Vec<_>>()
.join("\n\n");
tool.result = Some(pb::WebSearchResult {
result: Some(pb::web_search_result::Result::Success(
pb::WebSearchSuccess {
references: hits
.into_iter()
.map(|hit| pb::WebSearchReference {
title: hit.title,
url: hit.url,
chunk: hit.chunk,
})
.collect(),
},
)),
});
(output, false)
}
Err(error) => {
tool.result = Some(pb::WebSearchResult {
result: Some(pb::web_search_result::Result::Error(pb::WebSearchError {
error: error.clone(),
})),
});
(error, true)
}
};
ToolCompletion::from_rendered(call, pending.started_at_ms, output, is_error, rendered)
}
pub(crate) fn complete_web_fetch(
pending: PendingInteraction,
outcome: std::result::Result<FetchedPage, String>,
) -> Result<ToolCompletion> {
let call = &pending.call;
let requested_url = call
.arguments
.get("url")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::WebFetchToolCall(tool)) = rendered.tool.as_mut() else {
return Err(Error::Protocol(format!(
"tool {} is not WebFetch",
call.name
)));
};
let (output, is_error) = match outcome {
Ok(page) => {
let output = page.markdown.clone();
tool.result = Some(pb::WebFetchResult {
result: Some(pb::web_fetch_result::Result::Success(pb::WebFetchSuccess {
url: page.url,
markdown: page.markdown,
output_location: None,
})),
});
(output, false)
}
Err(error) => {
tool.result = Some(pb::WebFetchResult {
result: Some(pb::web_fetch_result::Result::Error(pb::WebFetchError {
url: requested_url.into(),
error: error.clone(),
})),
});
(error, true)
}
};
ToolCompletion::from_rendered(call, pending.started_at_ms, output, is_error, rendered)
}
fn ask_output(value: &pb::AskQuestionResult) -> Result<(String, bool)> {
use pb::ask_question_result::Result as R;
match value
.result
.as_ref()
.ok_or_else(|| missing("ask question"))?
{
R::Success(value) => Ok((
value
.answers
.iter()
.map(|answer| {
let value = if answer.freeform_text.is_empty() {
answer.selected_option_ids.join(", ")
} else {
answer.freeform_text.clone()
};
format!("{}: {value}", answer.question_id)
})
.collect::<Vec<_>>()
.join("\n"),
false,
)),
R::Error(value) => Ok((value.error_message.clone(), true)),
R::Rejected(value) => Ok((value.reason.clone(), true)),
R::Async(_) => Ok(("question is running asynchronously".into(), false)),
}
}
fn create_plan_output(value: &pb::CreatePlanResult) -> Result<(String, bool)> {
use pb::create_plan_result::Result as R;
match value
.result
.as_ref()
.ok_or_else(|| missing("create plan"))?
{
R::Success(_) => Ok((format!("plan created: {}", value.plan_uri), false)),
R::Error(value) => Ok((value.error.clone(), true)),
}
}
fn switch_mode_result(
value: &pb::SwitchModeRequestResponse,
) -> Result<(pb::SwitchModeResult, (String, bool))> {
use pb::{switch_mode_request_response::Result as Input, switch_mode_result::Result as Output};
match value
.result
.as_ref()
.ok_or_else(|| missing("switch mode"))?
{
Input::Approved(_) => Ok((
pb::SwitchModeResult {
result: Some(Output::Success(pb::SwitchModeSuccess::default())),
},
("mode switched".into(), false),
)),
Input::Rejected(value) => Ok((
pb::SwitchModeResult {
result: Some(Output::Rejected(pb::SwitchModeRejected {
reason: value.reason.clone(),
})),
},
(value.reason.clone(), true),
)),
}
}
fn missing(name: &str) -> Error {
Error::Protocol(format!("{name} returned no result"))
}
@@ -0,0 +1,178 @@
//! Converts server-local Tool completions into Tool results.
use serde_json::Value;
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec as interaction},
model::{ToolCall, ToolResult},
Error, Result,
};
use super::{now_ms, ToolCompletion};
const SUBAGENTS_DISABLED_REMINDER: &str = "<system_reminder>The user has disabled the subagent model. Please remind the user to enable it in Cursor Settings → Models → Explore Subagent Model.</system_reminder>";
pub(crate) fn local(call: &ToolCall, message_index: usize) -> Result<ToolCompletion> {
match normalized(&call.name).as_str() {
"todowrite" => todo_write(call),
"updatecurrentstep" => update_current_step(call, message_index),
_ => Err(Error::Protocol(format!("unsupported tool: {}", call.name))),
}
}
pub(crate) fn subagents_disabled(call: &ToolCall) -> Result<ToolCompletion> {
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::TaskToolCall(tool)) = rendered.tool.as_mut() else {
return Err(Error::Protocol("Task has no Cursor representation".into()));
};
tool.result = Some(pb::TaskResult {
result: Some(pb::task_result::Result::Error(pb::TaskError {
error: SUBAGENTS_DISABLED_REMINDER.into(),
})),
});
let tool = rendered
.tool
.ok_or_else(|| Error::Protocol("Task has no Cursor representation".into()))?;
Ok(ToolCompletion::new(
call,
now_ms(),
ToolResult {
call_id: call.call_id.clone(),
content: SUBAGENTS_DISABLED_REMINDER.into(),
is_error: true,
image: None,
},
tool,
))
}
fn todo_write(call: &ToolCall) -> Result<ToolCompletion> {
let todos = todo_items(&call.arguments);
let total_count = todos.len() as i32;
let was_merge = call
.arguments
.get("merge")
.and_then(Value::as_bool)
.unwrap_or(false);
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::UpdateTodosToolCall(tool)) = rendered.tool.as_mut() else {
return Err(Error::Protocol(
"TodoWrite has no Cursor representation".into(),
));
};
tool.result = Some(pb::UpdateTodosResult {
result: Some(pb::update_todos_result::Result::Success(
pb::UpdateTodosSuccess {
todos,
total_count,
was_merge,
},
)),
});
let tool = rendered
.tool
.ok_or_else(|| Error::Protocol("TodoWrite has no Cursor representation".into()))?;
Ok(ToolCompletion::new(
call,
now_ms(),
ToolResult {
call_id: call.call_id.clone(),
content: call.arguments.to_string(),
is_error: false,
image: None,
},
tool,
))
}
fn update_current_step(call: &ToolCall, message_index: usize) -> Result<ToolCompletion> {
let current_step = call
.arguments
.get("current_step")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string();
let mut rendered = interaction::render_tool_call(call, false)?;
let Some(pb::tool_call::Tool::CommunicateUpdateToolCall(tool)) = rendered.tool.as_mut() else {
return Err(Error::Protocol(
"UpdateCurrentStep has no Cursor representation".into(),
));
};
let message_index = u32::try_from(message_index)
.map_err(|_| Error::Protocol("Cursor message index space exhausted".into()))?;
tool.result = Some(pb::CommunicateUpdateResult {
result: Some(pb::communicate_update_result::Result::Success(
pb::CommunicateUpdateSuccess {
current_step: current_step.clone(),
message_index,
},
)),
});
let tool = rendered
.tool
.ok_or_else(|| Error::Protocol("UpdateCurrentStep has no Cursor representation".into()))?;
Ok(ToolCompletion::new(
call,
now_ms(),
ToolResult {
call_id: call.call_id.clone(),
content: serde_json::json!({
"success": {
"current_step": current_step,
"message_index": message_index,
}
})
.to_string(),
is_error: false,
image: None,
},
tool,
))
}
pub(crate) fn todo_items(arguments: &Value) -> Vec<pb::TodoItem> {
arguments
.get("todos")
.and_then(Value::as_array)
.into_iter()
.flatten()
.map(|todo| pb::TodoItem {
id: text(todo, "id"),
content: text(todo, "content"),
status: match todo
.get("status")
.and_then(Value::as_str)
.unwrap_or("pending")
{
"in_progress" => pb::TodoStatus::InProgress as i32,
"completed" => pb::TodoStatus::Completed as i32,
"cancelled" => pb::TodoStatus::Cancelled as i32,
_ => pb::TodoStatus::Pending as i32,
},
created_at: 0,
updated_at: 0,
dependencies: todo
.get("dependencies")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_string)
.collect(),
})
.collect()
}
fn text(value: &Value, name: &str) -> String {
value
.get(name)
.and_then(Value::as_str)
.unwrap_or_default()
.into()
}
fn normalized(name: &str) -> String {
name.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
@@ -0,0 +1,61 @@
//! Converts MCP completions into Tool results.
//! Canonical failures produced before an MCP request reaches the Cursor client.
use crate::{
cursor::{protocol::proto::agent::v1 as pb, tools::codec},
model::{ToolCall, ToolResult},
Result,
};
use super::{now_ms, ToolCompletion};
pub(crate) fn failure(call: &ToolCall, error: String) -> Result<ToolCompletion> {
let server = call
.arguments
.get("server")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let tool_name = call
.arguments
.get("toolName")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let arguments = call
.arguments
.get("arguments")
.and_then(serde_json::Value::as_object)
.map(codec::json_object_to_prost)
.unwrap_or_default();
Ok(ToolCompletion::new(
call,
now_ms(),
ToolResult {
call_id: call.call_id.clone(),
content: error.clone(),
is_error: true,
image: None,
},
pb::tool_call::Tool::McpToolCall(pb::McpToolCall {
args: Some(pb::McpArgs {
name: format!("{server}-{tool_name}"),
args: arguments,
tool_call_id: call.call_id.clone(),
provider_identifier: server.into(),
tool_name: tool_name.into(),
server_identifier: server.into(),
..Default::default()
}),
result: Some(pb::McpToolResult {
result: Some(pb::mcp_tool_result::Result::Error(pb::McpToolError {
error,
read_tool_def_reminder: String::new(),
})),
}),
description: call
.arguments
.get("description")
.and_then(serde_json::Value::as_str)
.map(str::to_string),
}),
))
}
@@ -0,0 +1,144 @@
//! Tracks MCP state required to build Tool results.
use serde_json::Value;
use crate::{cursor::protocol::proto::agent::v1 as pb, model::ToolResult, Error, Result};
use super::{prost_json, ToolCompletion};
use crate::cursor::tools::runtime::PendingExec;
pub(super) fn complete(
pending: PendingExec,
result: &pb::McpStateExecResult,
) -> Result<ToolCompletion> {
let call = &pending.call;
let server_filter = call.arguments.get("server").and_then(Value::as_str);
let tool_filter = call.arguments.get("toolName").and_then(Value::as_str);
if tool_filter.is_some() && server_filter.is_none() {
return Err(Error::Protocol(
"GetMcpTools toolName requires server".into(),
));
}
let pattern = call
.arguments
.get("pattern")
.and_then(Value::as_str)
.map(regex::Regex::new)
.transpose()
.map_err(|error| Error::Protocol(format!("invalid GetMcpTools pattern: {error}")))?;
let args = pb::GetMcpToolsArgs {
server: server_filter.map(str::to_string),
tool_name: tool_filter.map(str::to_string),
pattern: call
.arguments
.get("pattern")
.and_then(Value::as_str)
.map(str::to_string),
tool_call_id: call.call_id.clone(),
};
let (content, is_error, result) = match result
.result
.as_ref()
.ok_or_else(|| Error::Protocol("McpStateExecResult is missing result".into()))?
{
pb::mcp_state_exec_result::Result::Success(success) => {
let mut matches = Vec::new();
for server in success.servers.iter().filter(|server| {
server_filter.is_none_or(|value| value == server.server_identifier)
}) {
let status = server.status.as_deref().unwrap_or("unknown");
let server_matches_pattern = pattern
.as_ref()
.is_none_or(|pattern| pattern.is_match(&server.server_identifier));
let mut matched_tool = false;
for tool in &server.tools {
if tool_filter.is_some_and(|value| value != tool.tool_name)
|| (!server_matches_pattern
&& pattern
.as_ref()
.is_some_and(|pattern| !pattern.is_match(&tool.tool_name)))
{
continue;
}
matched_tool = true;
matches.push(serde_json::json!({
"server": server.server_identifier,
"serverName": server.server_name,
"serverStatus": status,
"toolName": tool.tool_name,
"description": tool.description,
"inputSchema": schema(tool),
}));
}
if !matched_tool && server_matches_pattern {
matches.push(serde_json::json!({
"server": server.server_identifier,
"serverName": server.server_name,
"serverStatus": status,
"tools": [],
}));
}
}
let mut content = serde_json::json!({ "tools": matches });
if server_filter.is_some() {
let instructions = success
.servers
.iter()
.filter(|server| {
server_filter.is_none_or(|value| value == server.server_identifier)
})
.flat_map(|server| &server.instructions)
.map(|value| value.instructions.as_str())
.filter(|value| !value.trim().is_empty())
.collect::<Vec<_>>();
if !instructions.is_empty() {
content["serverInstructions"] = serde_json::json!(instructions);
}
}
let content = serde_json::to_string_pretty(&content)?;
let wire = pb::get_mcp_tools_agent_result::Result::Success(pb::GetMcpToolsSuccess {
content: content.clone(),
output_file_path: None,
});
(content, false, wire)
}
pb::mcp_state_exec_result::Result::Error(error) => failure(&error.error),
pb::mcp_state_exec_result::Result::Rejected(rejected) => failure(&rejected.reason),
};
Ok(ToolCompletion::new(
call,
pending.started_at_ms,
ToolResult {
call_id: call.call_id.clone(),
content,
is_error,
image: None,
},
pb::tool_call::Tool::GetMcpToolsToolCall(pb::GetMcpToolsToolCall {
args: Some(args),
result: Some(pb::GetMcpToolsAgentResult {
result: Some(result),
}),
}),
))
}
fn failure(message: &str) -> (String, bool, pb::get_mcp_tools_agent_result::Result) {
(
message.into(),
true,
pb::get_mcp_tools_agent_result::Result::Error(pb::GetMcpToolsError {
error: message.into(),
}),
)
}
fn schema(tool: &pb::McpToolDefinition) -> Value {
let raw = tool.input_schema_json.clone().unwrap_or_else(|| {
tool.input_schema
.as_ref()
.map(prost_json)
.and_then(|value| serde_json::to_string(&value).ok())
.unwrap_or_else(|| "{}".into())
});
serde_json::from_str(&raw).unwrap_or(Value::String(raw))
}
@@ -0,0 +1,177 @@
//! Converts completed Tool work into canonical Tool results.
mod exec;
mod gate;
mod interaction;
mod local;
mod mcp;
mod mcp_state;
mod search;
use serde_json::Value;
use tokio::sync::mpsc;
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{ToolCall, ToolImageReference, ToolResult},
store::BlobId,
Error, Result,
};
use super::runtime::now_ms;
pub(crate) use exec::{edit_failure, from_exec};
pub(crate) use interaction::{complete_web_fetch, complete_web_search, from_interaction};
pub(crate) use local::{local, subagents_disabled, todo_items};
pub(crate) use mcp::failure as mcp_failure;
pub(crate) use search::complete as semble;
#[derive(Clone, Debug)]
pub struct ToolCompletion {
result: ToolResult,
tool_call: pb::ToolCall,
read_image: Option<ReadImage>,
}
#[derive(Clone, Debug)]
pub(crate) struct ReadImage {
pub(crate) data: Vec<u8>,
pub(crate) mime_type: String,
pub(crate) path: String,
}
impl ToolCompletion {
pub fn result(&self) -> &ToolResult {
&self.result
}
pub fn tool_call(&self) -> &pb::ToolCall {
&self.tool_call
}
pub(super) fn with_read_image(mut self, image: Option<ReadImage>) -> Self {
self.read_image = image;
self
}
pub(crate) fn take_read_image(&mut self) -> Option<ReadImage> {
self.read_image.take()
}
pub(crate) fn persist_read_image(&mut self, blob_id: &BlobId, image: &ReadImage) -> Result<()> {
self.result.content = format!("Read image file: {}", image.path);
self.result.image = Some(ToolImageReference {
blob_id: blob_id.to_base64(),
mime_type: image.mime_type.clone(),
path: image.path.clone(),
});
let Some(pb::tool_call::Tool::ReadToolCall(call)) = self.tool_call.tool.as_mut() else {
return Err(Error::Protocol(
"Read image completion has no Read tool state".into(),
));
};
let Some(pb::read_tool_result::Result::Success(success)) = call
.result
.as_mut()
.and_then(|result| result.result.as_mut())
else {
return Err(Error::Protocol(
"Read image completion has no success state".into(),
));
};
success.output = Some(pb::read_tool_success::Output::DataBlobId(
blob_id.as_bytes().to_vec(),
));
Ok(())
}
pub(crate) fn new(
call: &ToolCall,
started_at_ms: u64,
mut result: ToolResult,
mut tool: pb::tool_call::Tool,
) -> Self {
// Apply the model-visible size gate once, at the tool completion
// boundary. Canonical history and every provider projection then
// carry the same bounded result without reprocessing it.
gate::tool_completion(&call.name, &mut tool, &mut result.content);
Self {
result,
tool_call: pb::ToolCall {
tool_call_id: Some(call.call_id.clone()),
started_at_ms: Some(started_at_ms),
completed_at_ms: Some(now_ms()),
tool: Some(tool),
hook_additional_contexts: Vec::new(),
},
read_image: None,
}
}
pub(super) fn from_rendered(
call: &ToolCall,
started_at_ms: u64,
output: String,
is_error: bool,
rendered: pb::ToolCall,
) -> Result<Self> {
let tool = rendered.tool.ok_or_else(|| {
Error::Protocol(format!("tool {} has no Cursor representation", call.name))
})?;
Ok(Self::new(
call,
started_at_ms,
ToolResult {
call_id: call.call_id.clone(),
content: output,
is_error,
image: None,
},
tool,
))
}
}
#[derive(Clone)]
pub struct ToolResultSender(mpsc::UnboundedSender<Result<ToolCompletion>>);
pub struct ToolResultReceiver(mpsc::UnboundedReceiver<Result<ToolCompletion>>);
pub fn tool_result_channel() -> (ToolResultSender, ToolResultReceiver) {
let (sender, receiver) = mpsc::unbounded_channel();
(ToolResultSender(sender), ToolResultReceiver(receiver))
}
impl ToolResultSender {
pub fn send(&self, result: ToolCompletion) {
let _ = self.0.send(Ok(result));
}
pub fn send_error(&self, error: Error) {
let _ = self.0.send(Err(error));
}
}
impl ToolResultReceiver {
pub async fn recv(&mut self) -> Option<Result<ToolCompletion>> {
self.0.recv().await
}
}
pub(super) fn prost_json(value: &prost_types::Value) -> Value {
use prost_types::value::Kind;
match value.kind.as_ref() {
None | Some(Kind::NullValue(_)) => Value::Null,
Some(Kind::NumberValue(value)) => serde_json::Number::from_f64(*value)
.map(Value::Number)
.unwrap_or(Value::Null),
Some(Kind::StringValue(value)) => Value::String(value.clone()),
Some(Kind::BoolValue(value)) => Value::Bool(*value),
Some(Kind::StructValue(value)) => Value::Object(
value
.fields
.iter()
.map(|(key, value)| (key.clone(), prost_json(value)))
.collect(),
),
Some(Kind::ListValue(value)) => Value::Array(value.values.iter().map(prost_json).collect()),
}
}
@@ -0,0 +1,112 @@
//! Converts search completions into Tool results.
//! Cursor MCP-card rendering for direct Semble Agent tools.
use serde_json::Value;
use crate::{
cursor::protocol::proto::agent::v1 as pb,
model::{ToolCall, ToolResult},
Result,
};
use super::ToolCompletion;
const PROVIDER_IDENTIFIER: &str = "builtin-semble";
pub(crate) fn complete(
call: &ToolCall,
started_at_ms: u64,
output: std::result::Result<Value, String>,
) -> Result<ToolCompletion> {
use pb::{mcp_tool_result::Result as McpResult, tool_call::Tool};
let (tool_name, fallback_description) = match normalized(&call.name).as_str() {
"semblesearch" => ("search", "Search the codebase"),
"semblefindrelated" => ("find_related", "Find related code"),
_ => (call.name.as_str(), "Search the codebase"),
};
let description = call
.arguments
.get("description")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(fallback_description)
.to_owned();
let arguments = call
.arguments
.as_object()
.map(|arguments| {
let mut arguments = arguments.clone();
arguments.remove("description");
crate::cursor::tools::codec::json_object_to_prost(&arguments)
})
.unwrap_or_default();
let (content, is_error, result) = match output {
Ok(value) => {
let content = serde_json::to_string_pretty(&value)?;
let structured_content = value.as_object().map(|value| prost_types::Struct {
fields: crate::cursor::tools::codec::json_object_to_prost(value)
.into_iter()
.collect(),
});
(
content.clone(),
false,
McpResult::Success(pb::McpSuccess {
content: vec![pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text: content,
output_location: None,
},
)),
}],
is_error: false,
structured_content,
}),
)
}
Err(error) => (
error.clone(),
true,
McpResult::Error(pb::McpToolError {
error,
read_tool_def_reminder: String::new(),
}),
),
};
Ok(ToolCompletion::new(
call,
started_at_ms,
ToolResult {
call_id: call.call_id.clone(),
content,
is_error,
image: None,
},
Tool::McpToolCall(pb::McpToolCall {
args: Some(pb::McpArgs {
name: tool_name.into(),
args: arguments,
tool_call_id: call.call_id.clone(),
provider_identifier: PROVIDER_IDENTIFIER.into(),
tool_name: tool_name.into(),
server_identifier: PROVIDER_IDENTIFIER.into(),
..Default::default()
}),
result: Some(pb::McpToolResult {
result: Some(result),
}),
description: Some(description),
}),
))
}
fn normalized(value: &str) -> String {
value
.chars()
.filter(|character| character.is_ascii_alphanumeric())
.flat_map(char::to_lowercase)
.collect()
}
+145
View File
@@ -0,0 +1,145 @@
//! Provides the request-scoped input, subscription, and terminal interface.
use std::sync::{Arc, OnceLock};
use bytes::Bytes;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::{
cursor::{
conversation::TransportCommand,
protocol::{connect, proto::agent::v1 as pb},
services::observability::CursorTraceRecorder,
},
Error, Result,
};
use super::OutputHub;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransportParent {
pub request_id: String,
pub tool_call_id: String,
}
#[derive(Clone)]
pub struct TransportHandle {
request_id: String,
commands: mpsc::Sender<TransportCommand>,
output: Arc<OutputHub>,
conversation_id: Arc<OnceLock<String>>,
parent: Arc<OnceLock<TransportParent>>,
trace: Option<CursorTraceRecorder>,
disconnect: CancellationToken,
}
impl TransportHandle {
pub(crate) fn new(
request_id: String,
commands: mpsc::Sender<TransportCommand>,
output: Arc<OutputHub>,
trace: Option<CursorTraceRecorder>,
) -> Self {
Self {
request_id,
commands,
output,
conversation_id: Arc::new(OnceLock::new()),
parent: Arc::new(OnceLock::new()),
trace,
disconnect: CancellationToken::new(),
}
}
pub fn request_id(&self) -> &str {
&self.request_id
}
pub fn set_conversation_id(&self, conversation_id: &str) -> Result<()> {
if conversation_id.is_empty() {
return Err(Error::Protocol("Cursor conversation id is required".into()));
}
if self
.conversation_id
.get()
.is_some_and(|current| current != conversation_id)
{
return Err(Error::Protocol(format!(
"conflicting conversation ids for request {}",
self.request_id
)));
}
let _ = self.conversation_id.set(conversation_id.into());
Ok(())
}
pub fn conversation_id(&self) -> Option<&str> {
self.conversation_id.get().map(String::as_str)
}
pub fn set_parent(&self, parent: TransportParent) -> Result<()> {
if parent.request_id.is_empty() || parent.tool_call_id.is_empty() {
return Err(Error::Protocol("Cursor parent ids are required".into()));
}
if self.parent.get().is_some_and(|current| current != &parent) {
return Err(Error::Protocol(format!(
"conflicting parent ids for request {}",
self.request_id
)));
}
let _ = self.parent.set(parent);
Ok(())
}
pub fn parent(&self) -> Option<&TransportParent> {
self.parent.get()
}
pub async fn command(&self, command: TransportCommand) -> Result<()> {
self.commands
.send(command)
.await
.map_err(|_| Error::RunNotFound(self.request_id.clone()))
}
pub async fn disconnect(&self) {
let _ = self.commands.send(TransportCommand::Disconnect).await;
}
pub fn subscribe(&self) -> tokio::sync::mpsc::UnboundedReceiver<Bytes> {
self.output.subscribe()
}
pub fn emit_frame(&self, frame: Bytes) -> bool {
self.output.emit(frame)
}
pub fn emit(&self, message: &pb::AgentServerMessage) -> Result<()> {
if self.emit_frame(connect::encode_message(message)?) {
Ok(())
} else {
Err(Error::RunNotFound(self.request_id.clone()))
}
}
pub(crate) fn close_output(&self) -> bool {
self.output.close()
}
pub(crate) async fn wait_closed(&self) {
self.output.wait_closed().await;
}
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
self.trace.as_ref()
}
pub(crate) fn disconnect_token(&self) -> CancellationToken {
self.disconnect.clone()
}
pub(crate) fn mark_disconnected(&self) {
self.disconnect.cancel();
}
}
+30
View File
@@ -0,0 +1,30 @@
//! Orders Bidi messages by append sequence number.
use std::collections::BTreeMap;
pub struct OrderedInbox<T> {
next: i64,
pending: BTreeMap<i64, T>,
}
impl<T> OrderedInbox<T> {
pub fn starting_at(next: i64) -> Self {
Self {
next,
pending: BTreeMap::new(),
}
}
pub fn push(&mut self, seqno: i64, value: T) -> Vec<(i64, T)> {
if seqno < self.next || self.pending.contains_key(&seqno) {
return Vec::new();
}
self.pending.insert(seqno, value);
let mut ready = Vec::new();
while let Some(value) = self.pending.remove(&self.next) {
ready.push((self.next, value));
self.next = self.next.saturating_add(1);
}
ready
}
}
+11
View File
@@ -0,0 +1,11 @@
//! Owns request-ID-scoped upstream ordering and downstream transport.
mod handle;
mod inbox;
mod output;
mod registry;
pub use handle::*;
pub use inbox::*;
pub use output::*;
pub use registry::*;
+65
View File
@@ -0,0 +1,65 @@
//! Buffers, replays, broadcasts, and atomically closes downstream output.
use bytes::Bytes;
use tokio::sync::{mpsc, Notify};
#[derive(Default)]
pub struct OutputHub {
state: parking_lot::Mutex<OutputState>,
closed: Notify,
}
#[derive(Default)]
struct OutputState {
history: Vec<Bytes>,
subscribers: Vec<mpsc::UnboundedSender<Bytes>>,
closed: bool,
}
impl OutputHub {
pub fn emit(&self, frame: Bytes) -> bool {
let mut state = self.state.lock();
if state.closed {
return false;
}
state.history.push(frame.clone());
state
.subscribers
.retain(|subscriber| subscriber.send(frame.clone()).is_ok());
true
}
pub fn subscribe(&self) -> mpsc::UnboundedReceiver<Bytes> {
let (sender, receiver) = mpsc::unbounded_channel();
let mut state = self.state.lock();
for frame in &state.history {
let _ = sender.send(frame.clone());
}
if !state.closed {
state.subscribers.push(sender);
}
receiver
}
pub fn close(&self) -> bool {
let mut state = self.state.lock();
if state.closed {
return false;
}
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;
}
}
}
+140
View File
@@ -0,0 +1,140 @@
//! Maps request IDs to active transport handles.
use std::{collections::HashMap, sync::Arc};
use tokio::sync::{mpsc, Mutex, Notify};
use crate::{
cursor::{
conversation::ConversationRegistry, prompting::PromptCompiler,
services::observability::CursorTraceRecorder,
},
provider::Provider,
store::Store,
Result,
};
use super::{OutputHub, TransportHandle};
#[derive(Clone)]
pub struct TransportRegistry {
inner: Arc<RegistryInner>,
}
struct RegistryInner {
local: Mutex<HashMap<String, TransportHandle>>,
upstream: Mutex<HashMap<String, u64>>,
route_changed: Notify,
store: Store,
conversations: ConversationRegistry,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TransportRoute {
Local,
Upstream(u64),
}
impl TransportRegistry {
pub fn new(store: Store, provider: Arc<dyn Provider>, compiler: PromptCompiler) -> Self {
Self {
inner: Arc::new(RegistryInner {
local: Mutex::new(HashMap::new()),
upstream: Mutex::new(HashMap::new()),
route_changed: Notify::new(),
conversations: ConversationRegistry::new(store.clone(), provider, compiler),
store,
}),
}
}
pub fn store(&self) -> &Store {
&self.inner.store
}
pub fn conversations(&self) -> &ConversationRegistry {
&self.inner.conversations
}
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);
}
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());
drop(local);
self.inner.route_changed.notify_waiters();
self.inner
.conversations
.bind_transport(handle.clone(), receiver);
let registry = Arc::downgrade(&self.inner);
let request_id = request_id.to_string();
tokio::spawn(async move {
output.wait_closed().await;
if let Some(registry) = registry.upgrade() {
registry.local.lock().await.remove(&request_id);
}
});
Ok(handle)
}
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
self.inner.local.lock().await.get(request_id).cloned()
}
pub async fn mark_upstream(&self, request_id: &str) {
let mut upstream = self.inner.upstream.lock().await;
let generation = upstream.get(request_id).copied().unwrap_or_default() + 1;
upstream.insert(request_id.into(), generation);
drop(upstream);
self.inner.route_changed.notify_waiters();
}
pub async fn upstream(&self, request_id: &str) -> bool {
self.inner.upstream.lock().await.contains_key(request_id)
}
pub async fn wait_route(&self, request_id: &str) -> TransportRoute {
loop {
let changed = self.inner.route_changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
if self.inner.local.lock().await.contains_key(request_id) {
return TransportRoute::Local;
}
if let Some(generation) = self.inner.upstream.lock().await.get(request_id).copied() {
return TransportRoute::Upstream(generation);
}
changed.await;
}
}
pub fn finish_upstream(&self, request_id: String, generation: u64) {
let registry = self.clone();
tokio::spawn(async move {
let mut upstream = registry.inner.upstream.lock().await;
if upstream.get(&request_id) == Some(&generation) {
upstream.remove(&request_id);
}
});
}
pub async fn shutdown(&self) {
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;
}
}
}
+68
View File
@@ -0,0 +1,68 @@
//! Defines the server-wide error type and error conversions.
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
pub type Result<T, E = Error> = std::result::Result<T, E>;
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("configuration error: {0}")]
Config(String),
#[error("protocol error: {0}")]
Protocol(String),
#[error("provider error: {0}")]
Provider(String),
#[error("store error: {0}")]
Store(String),
#[error("run was cancelled")]
Cancelled,
#[error("run not found: {0}")]
RunNotFound(String),
#[error("database error: {0}")]
Database(#[from] sqlx::Error),
#[error("database migration error: {0}")]
Migration(#[from] sqlx::migrate::MigrateError),
#[error("http error: {0}")]
Http(#[from] reqwest::Error),
#[error("protobuf decode error: {0}")]
Decode(#[from] prost::DecodeError),
#[error("protobuf encode error: {0}")]
Encode(#[from] prost::EncodeError),
#[error("json error: {0}")]
Json(#[from] serde_json::Error),
#[error("io error: {0}")]
Io(#[from] std::io::Error),
}
impl IntoResponse for Error {
fn into_response(self) -> Response {
let status = match self {
Self::Config(_) | Self::Protocol(_) | Self::Decode(_) | Self::Json(_) => {
StatusCode::BAD_REQUEST
}
Self::RunNotFound(_) => StatusCode::NOT_FOUND,
Self::Provider(_) | Self::Http(_) => StatusCode::BAD_GATEWAY,
Self::Cancelled => StatusCode::CONFLICT,
Self::Store(_)
| Self::Database(_)
| Self::Migration(_)
| Self::Encode(_)
| Self::Io(_) => StatusCode::INTERNAL_SERVER_ERROR,
};
let code = match status {
StatusCode::BAD_REQUEST => "invalid_argument",
StatusCode::NOT_FOUND => "not_found",
StatusCode::CONFLICT => "aborted",
StatusCode::BAD_GATEWAY => "unavailable",
_ => "internal",
};
(
status,
Json(serde_json::json!({ "code": code, "message": self.to_string() })),
)
.into_response()
}
}

Some files were not shown because too many files have changed in this diff Show More