mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 03:56:45 +08:00
refactor: rebuild desktop app with Tauri
This commit is contained in:
@@ -0,0 +1,180 @@
|
||||
use std::{future::IntoFuture, net::SocketAddr, time::Duration};
|
||||
|
||||
use tokio::net::TcpListener;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
config::{Config, ConsoleSource},
|
||||
control,
|
||||
cursor::{
|
||||
handlers,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
CursorSessionRegistry,
|
||||
},
|
||||
harness::CursorHarness,
|
||||
provider::ProviderRouter,
|
||||
run::RunRegistry,
|
||||
store::Store,
|
||||
Result,
|
||||
};
|
||||
|
||||
pub struct App {
|
||||
config: Config,
|
||||
router: axum::Router,
|
||||
registry: CursorSessionRegistry,
|
||||
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 run_registry = RunRegistry::default();
|
||||
let registry = CursorSessionRegistry::new(store.clone(), provider, compiler, run_registry);
|
||||
let control = control::ControlService::new(store.clone())?;
|
||||
let harness = control.cursor_harness().clone();
|
||||
let mut router = handlers::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 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 => {} }
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn service_listener_falls_back_when_configured_port_is_busy() {
|
||||
let occupied = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let requested = occupied.local_addr().unwrap();
|
||||
let listener = bind_service_listener(requested, true).await.unwrap();
|
||||
assert_ne!(listener.local_addr().unwrap().port(), requested.port());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
use crate::model::{CanonicalMessage, RuntimeEvent, ToolResult};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum ClientCommand {
|
||||
ToolResult(ToolResult),
|
||||
RuntimeMessage(CanonicalMessage),
|
||||
RuntimeEvent(RuntimeEvent),
|
||||
ClientClosed { error: String },
|
||||
Cancel,
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use tokio::sync::oneshot;
|
||||
|
||||
use crate::model::{RevisionId, ToolCall, ToolRoundId, Usage};
|
||||
use crate::run::RunOutcome;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CommitCause {
|
||||
InitialMessages,
|
||||
ToolRoundStarted(ToolRoundId),
|
||||
ToolResult { call_id: String },
|
||||
FinalTurn,
|
||||
Compaction { summary: String },
|
||||
RuntimeEvent { event_id: String },
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum CommitBarrier {
|
||||
None,
|
||||
BeforeContinue(oneshot::Sender<std::result::Result<(), String>>),
|
||||
}
|
||||
|
||||
impl CommitBarrier {
|
||||
pub fn before_continue() -> (Self, oneshot::Receiver<std::result::Result<(), String>>) {
|
||||
let (sender, receiver) = oneshot::channel();
|
||||
(Self::BeforeContinue(sender), receiver)
|
||||
}
|
||||
|
||||
pub fn is_required(&self) -> bool {
|
||||
matches!(self, Self::BeforeContinue(_))
|
||||
}
|
||||
|
||||
pub fn complete(self, result: std::result::Result<(), String>) {
|
||||
if let Self::BeforeContinue(sender) = self {
|
||||
let _ = sender.send(result);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StateCommitted {
|
||||
pub revision_id: RevisionId,
|
||||
pub tool_round_version: u64,
|
||||
pub cause: CommitCause,
|
||||
pub barrier: CommitBarrier,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum ClientEvent {
|
||||
AutoCompactionStarted,
|
||||
AutoCompactionCompleted,
|
||||
TextStart,
|
||||
TextDelta(String),
|
||||
TextEnd,
|
||||
ThinkingStart,
|
||||
ThinkingDelta(String),
|
||||
ThinkingEnd {
|
||||
duration: Duration,
|
||||
},
|
||||
ToolCallStart {
|
||||
index: usize,
|
||||
call_id: String,
|
||||
name: String,
|
||||
model_call_id: String,
|
||||
},
|
||||
ToolCallArgumentsDelta {
|
||||
index: usize,
|
||||
delta: String,
|
||||
},
|
||||
ToolCallEnd {
|
||||
index: usize,
|
||||
},
|
||||
Usage(Usage),
|
||||
ExecuteToolRound {
|
||||
round_id: ToolRoundId,
|
||||
calls: Vec<ToolCall>,
|
||||
},
|
||||
StateCommitted(StateCommitted),
|
||||
Ended(RunOutcome),
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
mod command;
|
||||
mod event;
|
||||
mod session;
|
||||
|
||||
pub use command::*;
|
||||
pub use event::*;
|
||||
pub use session::*;
|
||||
@@ -0,0 +1,28 @@
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use super::{ClientCommand, ClientEvent};
|
||||
|
||||
pub struct ClientPort {
|
||||
pub commands: mpsc::Receiver<ClientCommand>,
|
||||
pub events: mpsc::Sender<ClientEvent>,
|
||||
}
|
||||
|
||||
pub struct ClientSession {
|
||||
pub commands: mpsc::Sender<ClientCommand>,
|
||||
pub events: mpsc::Receiver<ClientEvent>,
|
||||
}
|
||||
|
||||
pub fn session(capacity: usize) -> (ClientPort, ClientSession) {
|
||||
let (commands_tx, commands_rx) = mpsc::channel(capacity);
|
||||
let (events_tx, events_rx) = mpsc::channel(capacity);
|
||||
(
|
||||
ClientPort {
|
||||
commands: commands_rx,
|
||||
events: events_tx,
|
||||
},
|
||||
ClientSession {
|
||||
commands: commands_tx,
|
||||
events: events_rx,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
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";
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
#[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) => Duration::from_secs(300),
|
||||
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: Duration::from_secs(300),
|
||||
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)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn database_url_in(home_dir: &std::path::Path) -> Result<String> {
|
||||
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))?;
|
||||
|
||||
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}"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn managed_database_supports_home_paths_with_spaces() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let home_dir = directory.path().join("home with spaces");
|
||||
let database_url = database_url_in(&home_dir).unwrap();
|
||||
|
||||
let store = crate::store::Store::connect(&database_url).await.unwrap();
|
||||
drop(store);
|
||||
|
||||
let data_dir = home_dir.join(DATA_DIR_NAME);
|
||||
assert!(data_dir.join(DATABASE_FILE_NAME).is_file());
|
||||
|
||||
#[cfg(unix)]
|
||||
assert_eq!(
|
||||
fs::metadata(data_dir).unwrap().permissions().mode() & 0o777,
|
||||
0o700
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
//! 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 = "http://127.0.0.1:8080/api/v1/ads?placement=menu";
|
||||
pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
|
||||
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
||||
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
||||
pub(super) const DISABLED_AD_IDS_HEADER: &str = "disable-ad-ids";
|
||||
|
||||
#[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).await?))
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn filters_disabled_slots_without_limiting_menu_ads() {
|
||||
let slot = |id: &str, enabled| AdSlot {
|
||||
id: id.into(),
|
||||
enabled,
|
||||
placement: AdPlacement::Menu,
|
||||
target: AdTarget {
|
||||
title: id.into(),
|
||||
description: String::new(),
|
||||
image_url: "https://example.com/target.png".into(),
|
||||
},
|
||||
content: AdContent {
|
||||
title: id.into(),
|
||||
description: String::new(),
|
||||
image_url: "https://example.com/content.png".into(),
|
||||
details: Vec::new(),
|
||||
button: AdButton {
|
||||
label: "Open".into(),
|
||||
action: AdAction {
|
||||
action_type: AdActionType::OpenBrowser,
|
||||
url: "https://example.com".into(),
|
||||
},
|
||||
},
|
||||
},
|
||||
};
|
||||
let runtime = AdRuntime {
|
||||
slots: vec![
|
||||
slot("one", true),
|
||||
slot("disabled", false),
|
||||
slot("two", true),
|
||||
slot("three", true),
|
||||
slot("four", true),
|
||||
],
|
||||
}
|
||||
.into_menu_slots()
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
runtime
|
||||
.slots
|
||||
.iter()
|
||||
.map(|slot| slot.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["one", "two", "three", "four"]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
use axum::{extract::State, http::StatusCode, Json};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{
|
||||
model::{ProviderEndpoint, ProviderEndpointInput, ProviderModel, ProviderModelInput},
|
||||
Result,
|
||||
};
|
||||
|
||||
use super::{ControlService, DiscoveredModels};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum ProviderSelection {
|
||||
Existing { provider_id: i64 },
|
||||
New { input: ProviderEndpointInput },
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct CreateCursorModels {
|
||||
pub provider: ProviderSelection,
|
||||
pub models: Vec<ProviderModelInput>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct CreatedCursorModels {
|
||||
pub provider: ProviderEndpoint,
|
||||
pub models: Vec<ProviderModel>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct DiscoverCursorModels {
|
||||
pub provider: ProviderSelection,
|
||||
}
|
||||
|
||||
pub async fn create(
|
||||
State(service): State<ControlService>,
|
||||
Json(input): Json<CreateCursorModels>,
|
||||
) -> Result<(StatusCode, Json<CreatedCursorModels>)> {
|
||||
let (provider, models) = match input.provider {
|
||||
ProviderSelection::Existing { provider_id } => {
|
||||
let provider = service
|
||||
.providers()
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|provider| provider.provider_id == provider_id)
|
||||
.ok_or_else(|| crate::Error::RunNotFound(format!("provider {provider_id}")))?;
|
||||
let models = service.save_models(provider_id, &input.models).await?;
|
||||
(provider, models)
|
||||
}
|
||||
ProviderSelection::New { input: provider } => {
|
||||
service
|
||||
.create_provider_with_models(&provider, &input.models)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
Ok((
|
||||
StatusCode::CREATED,
|
||||
Json(CreatedCursorModels { provider, models }),
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn discover(
|
||||
State(service): State<ControlService>,
|
||||
Json(input): Json<DiscoverCursorModels>,
|
||||
) -> Result<Json<DiscoveredModels>> {
|
||||
match input.provider {
|
||||
ProviderSelection::Existing { provider_id } => {
|
||||
Ok(Json(service.discover_models(provider_id).await?))
|
||||
}
|
||||
ProviderSelection::New { input } => Ok(Json(service.discover_input(&input).await?)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
use axum::{extract::State, Json};
|
||||
|
||||
use crate::{
|
||||
harness::{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?,
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
mod ads;
|
||||
mod calls;
|
||||
mod cursor_models;
|
||||
mod harness;
|
||||
mod models;
|
||||
mod overview;
|
||||
mod providers;
|
||||
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, 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/providers",
|
||||
get(providers::list).post(providers::create),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/providers/{provider_id}",
|
||||
put(providers::update).delete(providers::remove),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/providers/{provider_id}/models/discover",
|
||||
post(models::discover),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/providers/{provider_id}/models",
|
||||
post(models::save),
|
||||
)
|
||||
.route("/__byok-api__/api/models", get(models::list))
|
||||
.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/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/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),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/harness/cursor/models",
|
||||
post(cursor_models::create),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/harness/cursor/models/discover",
|
||||
post(cursor_models::discover),
|
||||
)
|
||||
.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::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,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::{
|
||||
body::Body,
|
||||
http::{header, HeaderValue, Request},
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn control_routes_only_exist_below_the_reserved_namespace() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = crate::store::Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("control.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let router = api_router(ControlService::new(store).unwrap());
|
||||
|
||||
let response = router
|
||||
.clone()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/__byok-api__/api/providers")
|
||||
.header(header::ORIGIN, "tauri://localhost")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), axum::http::StatusCode::OK);
|
||||
assert_eq!(
|
||||
response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
|
||||
Some(&HeaderValue::from_static("tauri://localhost"))
|
||||
);
|
||||
|
||||
let response = router
|
||||
.clone()
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/__byok-api__/api/overview")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), axum::http::StatusCode::OK);
|
||||
|
||||
let response = router
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/api/providers")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), axum::http::StatusCode::NOT_FOUND);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn development_frontend_proxy_preserves_the_reserved_path_and_query() {
|
||||
let upstream = Router::new().route(
|
||||
"/__byok-api__/{*path}",
|
||||
get(|request: Request<Body>| async move {
|
||||
request.uri().path_and_query().unwrap().as_str().to_string()
|
||||
}),
|
||||
);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let task = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() });
|
||||
let router = frontend_proxy_router(format!("http://{address}").parse().unwrap());
|
||||
|
||||
let response = router
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/__byok-api__/src/index.tsx?direct=1")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
assert_eq!(body, "/__byok-api__/src/index.tsx?direct=1");
|
||||
task.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cors_only_allows_tauri_loopback_and_private_network_origins() {
|
||||
for origin in [
|
||||
"tauri://localhost",
|
||||
"http://tauri.localhost",
|
||||
"http://localhost:1420",
|
||||
"http://127.0.0.1:1420",
|
||||
"https://192.168.1.20:8443",
|
||||
"http://[::1]:1420",
|
||||
"http://[fd00::20]:1420",
|
||||
] {
|
||||
assert!(local_origin(&origin.parse().unwrap()), "{origin}");
|
||||
}
|
||||
for origin in [
|
||||
"https://example.com",
|
||||
"https://8.8.8.8",
|
||||
"https://localhost.example.com",
|
||||
"null",
|
||||
] {
|
||||
assert!(!local_origin(&origin.parse().unwrap()), "{origin}");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
use axum::{
|
||||
extract::{Path, State},
|
||||
http::StatusCode,
|
||||
Json,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{
|
||||
model::{ProviderModel, ProviderModelInput},
|
||||
Result,
|
||||
};
|
||||
|
||||
use super::{ControlService, DiscoveredModels};
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct SaveModels {
|
||||
pub models: Vec<ProviderModelInput>,
|
||||
}
|
||||
|
||||
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ProviderModel>>> {
|
||||
Ok(Json(service.models().await?))
|
||||
}
|
||||
|
||||
pub async fn save(
|
||||
State(service): State<ControlService>,
|
||||
Path(provider_id): Path<i64>,
|
||||
Json(input): Json<SaveModels>,
|
||||
) -> Result<(StatusCode, Json<Vec<ProviderModel>>)> {
|
||||
Ok((
|
||||
StatusCode::CREATED,
|
||||
Json(service.save_models(provider_id, &input.models).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<ProviderModelInput>,
|
||||
) -> Result<Json<ProviderModel>> {
|
||||
Ok(Json(service.update_model(&model_hash, &input).await?))
|
||||
}
|
||||
|
||||
pub async fn discover(
|
||||
State(service): State<ControlService>,
|
||||
Path(provider_id): Path<i64>,
|
||||
) -> Result<Json<DiscoveredModels>> {
|
||||
Ok(Json(service.discover_models(provider_id).await?))
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
//! 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>,
|
||||
provider_ids: 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(),
|
||||
range.provider_ids.as_deref(),
|
||||
)
|
||||
.await?,
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
use axum::{
|
||||
extract::{Path, State},
|
||||
http::StatusCode,
|
||||
Json,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
model::{ProviderEndpoint, ProviderEndpointInput},
|
||||
Result,
|
||||
};
|
||||
|
||||
use super::ControlService;
|
||||
|
||||
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ProviderEndpoint>>> {
|
||||
Ok(Json(service.providers().await?))
|
||||
}
|
||||
|
||||
pub async fn create(
|
||||
State(service): State<ControlService>,
|
||||
Json(input): Json<ProviderEndpointInput>,
|
||||
) -> Result<(StatusCode, Json<ProviderEndpoint>)> {
|
||||
Ok((
|
||||
StatusCode::CREATED,
|
||||
Json(service.create_provider(&input).await?),
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn update(
|
||||
State(service): State<ControlService>,
|
||||
Path(provider_id): Path<i64>,
|
||||
Json(input): Json<ProviderEndpointInput>,
|
||||
) -> Result<Json<ProviderEndpoint>> {
|
||||
Ok(Json(service.update_provider(provider_id, &input).await?))
|
||||
}
|
||||
|
||||
pub async fn remove(
|
||||
State(service): State<ControlService>,
|
||||
Path(provider_id): Path<i64>,
|
||||
) -> Result<StatusCode> {
|
||||
service.delete_provider(provider_id).await?;
|
||||
Ok(StatusCode::NO_CONTENT)
|
||||
}
|
||||
@@ -0,0 +1,559 @@
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use base64::{engine::general_purpose::STANDARD, Engine};
|
||||
use reqwest::header::{HeaderName, HeaderValue};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use url::Url;
|
||||
|
||||
use super::ads::{
|
||||
AdDismissalInput, AdRuntime, ADS_ENDPOINT, APP_VERSION_HEADER, DEVICE_ID_HEADER,
|
||||
DISABLED_AD_IDS_HEADER, OS_HEADER,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
harness::CursorHarness,
|
||||
model::{
|
||||
CursorRunTraceArtifact, CursorRunTraceSummary, LlmCallRequest, LlmCallResponseChunk,
|
||||
LlmCallSummary, Overview, ProviderEndpoint, ProviderEndpointInput, ProviderEndpointSecret,
|
||||
ProviderModel, ProviderModelInput, ProviderType,
|
||||
},
|
||||
store::{PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ControlService {
|
||||
store: Store,
|
||||
cursor_harness: CursorHarness,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct DiscoveredModels {
|
||||
pub models: Vec<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) -> Result<Self> {
|
||||
Ok(Self {
|
||||
cursor_harness: CursorHarness::new(store.clone())?,
|
||||
store,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn cursor_harness(&self) -> &CursorHarness {
|
||||
&self.cursor_harness
|
||||
}
|
||||
|
||||
pub(super) async fn ads(&self, disabled_ad_ids: Option<&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"))
|
||||
.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 providers(&self) -> Result<Vec<ProviderEndpoint>> {
|
||||
self.store.providers().await
|
||||
}
|
||||
|
||||
pub async fn create_provider(&self, input: &ProviderEndpointInput) -> Result<ProviderEndpoint> {
|
||||
self.store.create_provider(input).await
|
||||
}
|
||||
|
||||
pub async fn update_provider(
|
||||
&self,
|
||||
provider_id: i64,
|
||||
input: &ProviderEndpointInput,
|
||||
) -> Result<ProviderEndpoint> {
|
||||
self.store.update_provider(provider_id, input).await
|
||||
}
|
||||
|
||||
pub async fn delete_provider(&self, provider_id: i64) -> Result<()> {
|
||||
self.store.delete_provider(provider_id).await
|
||||
}
|
||||
|
||||
pub async fn models(&self) -> Result<Vec<ProviderModel>> {
|
||||
self.store.provider_models(false).await
|
||||
}
|
||||
|
||||
pub async fn overview(
|
||||
&self,
|
||||
start_ms: Option<i64>,
|
||||
end_ms: Option<i64>,
|
||||
model_hashes: Option<&str>,
|
||||
provider_ids: Option<&str>,
|
||||
) -> Result<Overview> {
|
||||
self.store
|
||||
.overview(start_ms, end_ms, model_hashes, provider_ids)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn save_models(
|
||||
&self,
|
||||
provider_id: i64,
|
||||
models: &[ProviderModelInput],
|
||||
) -> Result<Vec<ProviderModel>> {
|
||||
self.store.save_provider_models(provider_id, models).await
|
||||
}
|
||||
|
||||
pub async fn delete_model(&self, model_hash: &str) -> Result<()> {
|
||||
self.store.delete_provider_model(model_hash).await
|
||||
}
|
||||
|
||||
pub async fn update_model(
|
||||
&self,
|
||||
model_hash: &str,
|
||||
input: &ProviderModelInput,
|
||||
) -> Result<ProviderModel> {
|
||||
self.store.update_provider_model(model_hash, input).await
|
||||
}
|
||||
|
||||
pub async fn create_provider_with_models(
|
||||
&self,
|
||||
provider: &ProviderEndpointInput,
|
||||
models: &[ProviderModelInput],
|
||||
) -> Result<(ProviderEndpoint, Vec<ProviderModel>)> {
|
||||
self.store
|
||||
.create_provider_with_models(provider, models)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn discover_input(&self, input: &ProviderEndpointInput) -> Result<DiscoveredModels> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let endpoint = ProviderEndpoint {
|
||||
provider_id: 0,
|
||||
name: input.name.clone(),
|
||||
provider_type: input.provider_type,
|
||||
base_url: crate::model::normalize_base_url(&input.base_url)?,
|
||||
has_api_key: input
|
||||
.api_key
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.is_empty()),
|
||||
custom_headers: input.custom_headers.clone(),
|
||||
extra_params: input.extra_params.clone(),
|
||||
created_at_ms: 0,
|
||||
updated_at_ms: 0,
|
||||
};
|
||||
let secret = ProviderEndpointSecret {
|
||||
endpoint,
|
||||
api_key: input.api_key.clone().unwrap_or_default(),
|
||||
custom_headers: input.custom_headers.clone(),
|
||||
};
|
||||
let mut models = match input.provider_type {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
|
||||
openai_models(&client, &secret).await?
|
||||
}
|
||||
ProviderType::Anthropic => anthropic_models(&client, &secret).await?,
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
Ok(DiscoveredModels { models })
|
||||
}
|
||||
|
||||
pub async fn discover_models(&self, provider_id: i64) -> Result<DiscoveredModels> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let provider = self
|
||||
.store
|
||||
.provider(provider_id)
|
||||
.await?
|
||||
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
|
||||
let mut models = match provider.endpoint.provider_type {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
|
||||
openai_models(&client, &provider).await?
|
||||
}
|
||||
ProviderType::Anthropic => anthropic_models(&client, &provider).await?,
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
Ok(DiscoveredModels { models })
|
||||
}
|
||||
|
||||
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 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
|
||||
}
|
||||
}
|
||||
|
||||
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,
|
||||
finished_at_ms: trace.finished_at_ms,
|
||||
queue_ms: None,
|
||||
ttfb_ms: ttfb,
|
||||
ttft_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 openai_models(
|
||||
client: &reqwest::Client,
|
||||
provider: &ProviderEndpointSecret,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut request = client.get(format!("{}/models", provider.endpoint.base_url));
|
||||
if !provider.api_key.is_empty() {
|
||||
request = request.bearer_auth(&provider.api_key);
|
||||
}
|
||||
let response = apply_custom_headers(request, &provider.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,
|
||||
provider: &ProviderEndpointSecret,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut after_id = None::<String>;
|
||||
let mut found = BTreeSet::new();
|
||||
loop {
|
||||
let mut request = client
|
||||
.get(format!("{}/models", provider.endpoint.base_url))
|
||||
.query(&[("limit", "100")])
|
||||
.header("anthropic-version", "2023-06-01");
|
||||
if !provider.api_key.is_empty() {
|
||||
request = request.header("x-api-key", &provider.api_key);
|
||||
}
|
||||
if let Some(after_id) = &after_id {
|
||||
request = request.query(&[("after_id", after_id)]);
|
||||
}
|
||||
let response = apply_custom_headers(request, &provider.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 apply_custom_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 {
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
use crate::Result;
|
||||
use axum::{extract::State, Json};
|
||||
|
||||
use crate::store::{PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage};
|
||||
|
||||
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>,
|
||||
) -> Result<Json<StatisticsStorage>> {
|
||||
Ok(Json(service.clear_statistics_storage().await?))
|
||||
}
|
||||
|
||||
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?))
|
||||
}
|
||||
@@ -0,0 +1,467 @@
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
extract::Extension,
|
||||
http::{header, Request, Response},
|
||||
};
|
||||
use prost::Message;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{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
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::{
|
||||
body::to_bytes,
|
||||
http::StatusCode,
|
||||
routing::{get, post},
|
||||
Extension, Router,
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn app(upstream: Router) -> (Router, tokio::task::JoinHandle<()>) {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() });
|
||||
let proxy = proxy::CursorProxy::for_upstream(&format!("http://{address}")).unwrap();
|
||||
let app = Router::new()
|
||||
.route("/auth/full_stripe_profile", get(stripe_profile))
|
||||
.route("/aiserver.v1.DashboardService/GetMe", post(get_me))
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetCurrentPeriodUsage",
|
||||
post(current_period_usage),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
|
||||
post(usage_limit_status),
|
||||
)
|
||||
.layer(Extension(proxy));
|
||||
(app, server)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_upstream_profile_and_overlays_ultra_membership() {
|
||||
let upstream = Router::new().route(
|
||||
"/auth/full_stripe_profile",
|
||||
get(|| async {
|
||||
axum::Json(serde_json::json!({
|
||||
"membershipType": "pro",
|
||||
"subscriptionStatus": "inactive",
|
||||
"paymentId": "upstream-payment"
|
||||
}))
|
||||
}),
|
||||
);
|
||||
let (app, server) = app(upstream).await;
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::get("/auth/full_stripe_profile")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
let profile: Value = serde_json::from_slice(&body).unwrap();
|
||||
assert_eq!(profile["membershipType"], "ultra");
|
||||
assert_eq!(profile["paymentId"], "upstream-payment");
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upstream_error_uses_local_identity_without_reading_authorization() {
|
||||
let upstream = Router::new().route(
|
||||
"/aiserver.v1.DashboardService/GetMe",
|
||||
post(|| async { StatusCode::UNAUTHORIZED }),
|
||||
);
|
||||
let (app, server) = app(upstream).await;
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::post("/aiserver.v1.DashboardService/GetMe")
|
||||
.header(header::AUTHORIZATION, "Bearer ignored")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
let identity = GetMeResponse::decode(body).unwrap();
|
||||
assert_eq!(identity.auth_id, LOCAL_AUTH_ID);
|
||||
assert_eq!(identity.email.as_deref(), Some(LOCAL_EMAIL));
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stripe_error_uses_the_complete_local_ultra_profile() {
|
||||
let upstream = Router::new().route(
|
||||
"/auth/full_stripe_profile",
|
||||
get(|| async { StatusCode::SERVICE_UNAVAILABLE }),
|
||||
);
|
||||
let (app, server) = app(upstream).await;
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::get("/auth/full_stripe_profile")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
let profile: Value = serde_json::from_slice(&body).unwrap();
|
||||
assert_eq!(profile["membershipType"], "ultra");
|
||||
assert_eq!(profile["paymentId"], LOCAL_AUTH_ID);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn current_period_usage_is_a_local_unused_ultra_allowance() {
|
||||
let (app, server) = app(Router::new()).await;
|
||||
let before = chrono::Utc::now().timestamp_millis();
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::post("/aiserver.v1.DashboardService/GetCurrentPeriodUsage")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
let usage = GetCurrentPeriodUsageResponse::decode(body).unwrap();
|
||||
let plan = usage.plan_usage.unwrap();
|
||||
assert_eq!(plan.total_spend, 0);
|
||||
assert_eq!(plan.limit, LOCAL_ULTRA_PLAN_INCLUDED_CENTS);
|
||||
assert_eq!(plan.remaining, LOCAL_ULTRA_PLAN_INCLUDED_CENTS);
|
||||
assert_eq!(usage.display_message, "Ultra plan active");
|
||||
assert!(usage.billing_cycle_start < before);
|
||||
assert!(usage.billing_cycle_end > before);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn usage_limit_status_is_local_and_unrestricted() {
|
||||
let (app, server) = app(Router::new()).await;
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::post("/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
let response = GetUsageLimitStatusAndActiveGrantsResponse::decode(body).unwrap();
|
||||
let policy = response.usage_limit_policy_status.unwrap();
|
||||
assert!(!policy.is_in_slow_pool);
|
||||
assert!(policy.can_configure_spend_limit);
|
||||
assert!(!policy.has_pending_request);
|
||||
assert!(policy.allowed_model_ids.is_empty());
|
||||
assert!(policy.allowed_model_tags.is_empty());
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,367 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::{
|
||||
cursor::prompting::PromptCompiler,
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer,
|
||||
checkpoint::CheckpointBuilder,
|
||||
context_sync::RequestContextSynchronizer,
|
||||
proto::agent::v1 as pb,
|
||||
request,
|
||||
session::CursorSession,
|
||||
tools::{
|
||||
codec, result::tool_result_channel, runtime::CursorToolRuntime, ClientToolEvent,
|
||||
ToolDispatcher,
|
||||
},
|
||||
},
|
||||
provider::Provider,
|
||||
run::{RunActor, RunRegistry},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::{inbox::OrderedInbox, CursorCommand, CursorSessionHandle};
|
||||
|
||||
pub struct CursorActor;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct RunDependencies {
|
||||
pub store: Store,
|
||||
pub provider: Arc<dyn Provider>,
|
||||
pub compiler: PromptCompiler,
|
||||
pub run_registry: RunRegistry,
|
||||
}
|
||||
|
||||
impl CursorActor {
|
||||
pub(crate) fn spawn(
|
||||
handle: CursorSessionHandle,
|
||||
mut receiver: mpsc::Receiver<CursorCommand>,
|
||||
dependencies: RunDependencies,
|
||||
blob_sync: BlobSynchronizer,
|
||||
next_append_seqno: i64,
|
||||
) {
|
||||
tokio::spawn(async move {
|
||||
let mut inbox = OrderedInbox::starting_at(next_append_seqno);
|
||||
let (results_tx, results_rx) = tool_result_channel();
|
||||
let (runtime_actions_tx, runtime_actions_rx) = mpsc::unbounded_channel();
|
||||
let tool_runtime = CursorToolRuntime::default();
|
||||
let context_sync =
|
||||
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
|
||||
let tools = ToolDispatcher::with_results(tool_runtime.clone(), results_tx.clone());
|
||||
let mut run_resources = Some((results_rx, runtime_actions_rx, dependencies));
|
||||
loop {
|
||||
let command = match receiver.recv().await {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.cancel();
|
||||
break;
|
||||
}
|
||||
};
|
||||
match command {
|
||||
CursorCommand::Abort => {
|
||||
handle.cancel();
|
||||
}
|
||||
CursorCommand::Finished => {
|
||||
break;
|
||||
}
|
||||
CursorCommand::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((results, runtime_actions, dependencies)) =
|
||||
run_resources.take()
|
||||
{
|
||||
let handle = handle.clone();
|
||||
let blob_sync = blob_sync.clone();
|
||||
let context_sync = context_sync.clone();
|
||||
let tools = tools.clone();
|
||||
let tool_runtime = tool_runtime.clone();
|
||||
tokio::spawn(async move {
|
||||
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 parent = handle.parent().map(|parent| {
|
||||
(
|
||||
crate::model::RunId::new(&parent.run_id),
|
||||
parent.tool_call_id.clone(),
|
||||
)
|
||||
});
|
||||
let prepared = request::prepare(
|
||||
handle.request_id(),
|
||||
&request,
|
||||
parent,
|
||||
request::PrepareDependencies {
|
||||
compiler: &dependencies.compiler,
|
||||
store: &dependencies.store,
|
||||
checkpoint: &checkpoint,
|
||||
blob_sync: &blob_sync,
|
||||
context_sync: &context_sync,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
let (prepared, context) = match prepared {
|
||||
Ok(prepared) => prepared,
|
||||
Err(error) => {
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"failed to prepare Cursor Run"
|
||||
);
|
||||
let _ = crate::cursor::lifecycle::fail(
|
||||
&handle, &error,
|
||||
);
|
||||
let _ = handle
|
||||
.command(CursorCommand::Finished)
|
||||
.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(),
|
||||
);
|
||||
let cancellation = handle.cancellation();
|
||||
let (port, core) = crate::client::session(256);
|
||||
let actor = RunActor::new(
|
||||
dependencies.store.clone(),
|
||||
dependencies.provider,
|
||||
dependencies.run_registry,
|
||||
);
|
||||
let core_run =
|
||||
actor.spawn(prepared, port, cancellation).await;
|
||||
let session = CursorSession::new(
|
||||
handle.clone(),
|
||||
dependencies.store,
|
||||
context,
|
||||
core,
|
||||
super::session::CursorSessionRuntime {
|
||||
tools,
|
||||
results,
|
||||
runtime_actions,
|
||||
compiler: dependencies.compiler,
|
||||
blob_sync,
|
||||
checkpoint,
|
||||
tool_runtime,
|
||||
},
|
||||
);
|
||||
if let Err(error) = session.run().await {
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"Cursor session failed"
|
||||
);
|
||||
handle.cancel();
|
||||
let _ = crate::cursor::lifecycle::fail(
|
||||
&handle, &error,
|
||||
);
|
||||
}
|
||||
let _ = core_run.await;
|
||||
let _ =
|
||||
handle.command(CursorCommand::Finished).await;
|
||||
});
|
||||
} else {
|
||||
let error = crate::Error::Protocol(format!(
|
||||
"duplicate RunRequest for request_id: {}",
|
||||
handle.request_id()
|
||||
));
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"rejected duplicate Cursor RunRequest"
|
||||
);
|
||||
results_tx.send_error(error);
|
||||
}
|
||||
}
|
||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
message,
|
||||
)) => {
|
||||
if context_sync.handle_client(&message).await {
|
||||
continue;
|
||||
}
|
||||
match codec::client_event(&message, &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)) => {
|
||||
results_tx.send(*result)
|
||||
}
|
||||
Ok(codec::ClientExecEvent::Pending) => {}
|
||||
Err(error) => results_tx.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;
|
||||
}
|
||||
if tool_runtime.take_exec(close.id).await.is_some()
|
||||
{
|
||||
results_tx.send_error(crate::Error::Protocol(format!(
|
||||
"Exec stream closed before result for id: {}",
|
||||
close.id
|
||||
)));
|
||||
}
|
||||
}
|
||||
Some(Message::Throw(throw)) => {
|
||||
if context_sync
|
||||
.handle_throw(
|
||||
throw.id,
|
||||
format!(
|
||||
"Cursor request context failed: {}",
|
||||
throw.error
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
continue;
|
||||
}
|
||||
match tool_runtime.take_exec(throw.id).await {
|
||||
Some(pending) => results_tx.send_error(
|
||||
crate::Error::Protocol(format!(
|
||||
"Exec {} failed: {}",
|
||||
pending.call.call_id, throw.error
|
||||
)),
|
||||
),
|
||||
None => results_tx.send_error(
|
||||
crate::Error::Protocol(format!(
|
||||
"unknown ExecClientThrow id: {}",
|
||||
throw.id
|
||||
)),
|
||||
),
|
||||
}
|
||||
}
|
||||
Some(Message::Heartbeat(_)) | None => {}
|
||||
}
|
||||
}
|
||||
Some(
|
||||
pb::agent_client_message::Message::InteractionResponse(
|
||||
message,
|
||||
),
|
||||
) => match tools.interaction_response(&message).await {
|
||||
Ok(ClientToolEvent::Completed(completion)) => {
|
||||
results_tx.send(*completion)
|
||||
}
|
||||
Ok(ClientToolEvent::Pending) => {}
|
||||
Err(error) => results_tx.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. request::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 request::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, CancelSubagentAction,
|
||||
// BackgroundShellAction, BackgroundSubagentAction,
|
||||
// SubscriptionNotificationAction and GoalContinuationAction.
|
||||
// CancelSubagentAction must not start an LLM; 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(_),
|
||||
)
|
||||
| Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||
handle.cancel();
|
||||
}
|
||||
Some(
|
||||
pb::conversation_action::Action::InjectContextAction(
|
||||
action,
|
||||
),
|
||||
) => {
|
||||
if runtime_actions_tx.send(action).is_err() {
|
||||
results_tx.send_error(crate::Error::Protocol(
|
||||
"InjectContextAction arrived without an active Run"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
Some(action) => {
|
||||
results_tx.send_error(crate::Error::Protocol(format!(
|
||||
"unsupported runtime ConversationAction: {}",
|
||||
runtime_action_name(&action)
|
||||
)));
|
||||
}
|
||||
None => results_tx.send_error(crate::Error::Protocol(
|
||||
"runtime ConversationAction has no action".into(),
|
||||
)),
|
||||
},
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
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",
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
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::{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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn overlays_retry_gate_without_losing_upstream_config() {
|
||||
let mut config = json!({
|
||||
"feature_gates": {
|
||||
"upstream_gate": { "name": "upstream_gate", "value": true }
|
||||
},
|
||||
"dynamic_configs": { "kept": { "value": 1 } }
|
||||
});
|
||||
|
||||
enable_agent_retries(&mut config).unwrap();
|
||||
|
||||
assert_eq!(config["feature_gates"][AGENT_RETRIES_GATE]["value"], true);
|
||||
assert_eq!(config["feature_gates"]["upstream_gate"]["value"], true);
|
||||
assert_eq!(config["dynamic_configs"]["kept"]["value"], 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uses_the_hash_algorithm_declared_by_upstream() {
|
||||
let mut config = json!({
|
||||
"hash_used": "djb2",
|
||||
"feature_gates": {}
|
||||
});
|
||||
|
||||
enable_agent_retries(&mut config).unwrap();
|
||||
|
||||
let key = djb2(AGENT_RETRIES_GATE);
|
||||
assert_eq!(config["feature_gates"][&key]["name"], key);
|
||||
assert_eq!(config["feature_gates"][&key]["value"], true);
|
||||
assert!(config["feature_gates"].get(AGENT_RETRIES_GATE).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn patches_raw_and_connect_framed_responses() {
|
||||
for framed in [false, true] {
|
||||
let message = BootstrapStatsigResponse {
|
||||
config: json!({ "feature_gates": {} }).to_string(),
|
||||
generated_at_ms: 123,
|
||||
};
|
||||
let body = encode_unary(&message, framed);
|
||||
let buffered = proxy::BufferedResponse {
|
||||
status: StatusCode::OK,
|
||||
headers: Default::default(),
|
||||
body,
|
||||
};
|
||||
|
||||
let response = patch_upstream(buffered).unwrap();
|
||||
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.unwrap();
|
||||
let (_, payload) = unary_payload(&body).unwrap();
|
||||
let patched = BootstrapStatsigResponse::decode(payload).unwrap();
|
||||
let config: Value = serde_json::from_str(&patched.config).unwrap();
|
||||
assert_eq!(config["feature_gates"][AGENT_RETRIES_GATE]["value"], true);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,221 @@
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
cursor::{CursorCommand, CursorParent, CursorSessionRegistry},
|
||||
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 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: &CursorSessionRegistry,
|
||||
request: DecodedAppend,
|
||||
parent: Option<CursorParent>,
|
||||
) -> Result<ai::BidiAppendResponse> {
|
||||
let handle = registry.get_or_create(&request.request_id).await?;
|
||||
if let Some(parent) = parent {
|
||||
handle.set_parent(parent)?;
|
||||
}
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: request.seqno,
|
||||
message: Box::new(request.message),
|
||||
})
|
||||
.await?;
|
||||
Ok(ai::BidiAppendResponse {})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn encoded(run: agent::AgentRunRequest) -> ai::BidiAppendRequest {
|
||||
let message = agent::AgentClientMessage {
|
||||
message: Some(agent::agent_client_message::Message::RunRequest(run)),
|
||||
};
|
||||
ai::BidiAppendRequest {
|
||||
data: hex::encode(message.encode_to_vec()),
|
||||
request_id: Some(ai::BidiRequestId {
|
||||
request_id: "request".into(),
|
||||
}),
|
||||
append_seqno: 1,
|
||||
data_binary: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn route_model_uses_requested_model_id() {
|
||||
let decoded = decode(&encoded(agent::AgentRunRequest {
|
||||
requested_model: Some(agent::RequestedModel {
|
||||
model_id: "33ceed20".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(decoded.model_id(), Some("33ceed20"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn route_model_uses_legacy_model_details_when_needed() {
|
||||
let decoded = decode(&encoded(agent::AgentRunRequest {
|
||||
model_details: Some(agent::ModelDetails {
|
||||
model_id: "grok-4.6".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(decoded.model_id(), Some("grok-4.6"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
sync::{
|
||||
atomic::{AtomicU32, Ordering},
|
||||
Arc,
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use tokio::sync::{oneshot, Mutex};
|
||||
|
||||
use crate::{
|
||||
cursor::observability::CursorTraceRecorder,
|
||||
cursor::proto::agent::v1 as pb,
|
||||
cursor::CursorSessionHandle,
|
||||
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: CursorSessionHandle,
|
||||
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: CursorSessionHandle) -> 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.cancellation();
|
||||
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(15)) => 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.cancellation();
|
||||
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(15)) => 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(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{prompting::fold_derived_state, 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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::{Origin, Role, ToolCallContent, ToolResultContent};
|
||||
|
||||
#[test]
|
||||
fn update_current_step_is_folded_from_canonical_messages() {
|
||||
let messages = vec![
|
||||
CanonicalMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
origin: Origin::Assistant,
|
||||
content: MessageContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: String::new(),
|
||||
tool_round_id: Some("round".into()),
|
||||
replay_state: None,
|
||||
tool_calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
name: "UpdateCurrentStep".into(),
|
||||
arguments: serde_json::json!({
|
||||
"current_step": "Inspecting protocol",
|
||||
"final_summary": "Protocol verified.",
|
||||
"completed_subtitle": "Verified protocol flow"
|
||||
}),
|
||||
}],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
},
|
||||
CanonicalMessage {
|
||||
message_id: "result".into(),
|
||||
role: Role::Tool,
|
||||
origin: Origin::Tool,
|
||||
content: MessageContent::ToolResult(ToolResultContent {
|
||||
call_id: "call".into(),
|
||||
name: "UpdateCurrentStep".into(),
|
||||
content: serde_json::json!({
|
||||
"success": {"current_step": "Inspecting protocol", "message_index": 3}
|
||||
})
|
||||
.to_string(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
runtime_event_id: None,
|
||||
},
|
||||
];
|
||||
let state = update_current_step_state(&messages).unwrap();
|
||||
assert_eq!(state.history[0].step, "Inspecting protocol");
|
||||
assert_eq!(state.history[0].message_index, 3);
|
||||
assert_eq!(state.final_summary.as_deref(), Some("Protocol verified."));
|
||||
assert_eq!(
|
||||
state.completed_subtitle.as_deref(),
|
||||
Some("Verified protocol flow")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
mod derived;
|
||||
mod recovery;
|
||||
mod roots;
|
||||
mod summary;
|
||||
mod turns;
|
||||
pub(crate) mod worker;
|
||||
|
||||
use std::collections::HashSet;
|
||||
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer, presentation::PresentationDelta, projection,
|
||||
proto::agent::v1 as pb, CursorSessionHandle,
|
||||
},
|
||||
model::{CanonicalMessage, ToolCall, ToolDefinition, ToolRoundAssistant},
|
||||
store::Store,
|
||||
Result,
|
||||
};
|
||||
|
||||
use roots::RootFrontier;
|
||||
use turns::TurnFrontier;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CheckpointBuilder {
|
||||
store: Store,
|
||||
sync: BlobSynchronizer,
|
||||
parent_tool_call_id: Option<String>,
|
||||
base: pb::ConversationStateStructure,
|
||||
model: String,
|
||||
max_context_tokens: Option<u64>,
|
||||
instructions: String,
|
||||
tool_definitions: Vec<ToolDefinition>,
|
||||
allowed_tools: Vec<String>,
|
||||
dynamic_tools: HashSet<String>,
|
||||
turn_user: Option<pb::UserMessage>,
|
||||
roots: Option<RootFrontier>,
|
||||
turn: Option<TurnFrontier>,
|
||||
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: &PresentationDelta,
|
||||
) -> 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: &PresentationDelta,
|
||||
) -> Result<pb::ConversationStateStructure> {
|
||||
let pending = projection::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: &PresentationDelta,
|
||||
) -> Result<pb::ConversationStateStructure> {
|
||||
let pending = projection::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: &PresentationDelta,
|
||||
) -> 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::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: &PresentationDelta) {
|
||||
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: &CursorSessionHandle,
|
||||
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))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::context_limit;
|
||||
|
||||
#[test]
|
||||
fn selected_context_replaces_checkpoint_context() {
|
||||
assert_eq!(context_limit(Some(800_000), Some(200_000)), Some(800_000));
|
||||
assert_eq!(context_limit(None, Some(200_000)), Some(200_000));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
use crate::{
|
||||
cursor::{projection, 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(projection::decode(
|
||||
&data,
|
||||
format!("cursor-root:{}:{ordinal}", id.to_base64()),
|
||||
)?);
|
||||
}
|
||||
Ok(messages)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
use crate::{cursor::projection, 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 = projection::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 = projection::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())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn a_new_prompt_replaces_only_the_system_root() {
|
||||
let previous = b"previous prompt".to_vec();
|
||||
let current = b"current prompt".to_vec();
|
||||
let roots = RootFrontier {
|
||||
ids: vec![BlobId::digest(&previous), BlobId::digest(b"user")],
|
||||
generated: Vec::new(),
|
||||
base_count: 2,
|
||||
};
|
||||
assert_eq!(
|
||||
changed_system_root(&roots, &[current.clone(), b"user".to_vec()]),
|
||||
Some(current)
|
||||
);
|
||||
assert_eq!(
|
||||
changed_system_root(&roots, &[previous, b"user".to_vec()]),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{presentation::PresentationDelta, 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: &PresentationDelta,
|
||||
) -> 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::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())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{presentation::PresentationDelta, 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: &PresentationDelta,
|
||||
) -> 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(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::{
|
||||
cursor::{presentation::PresentationDelta, proto::agent::v1 as pb, CursorSessionHandle},
|
||||
model::{RevisionId, ToolRoundId},
|
||||
store::Store,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::CheckpointBuilder;
|
||||
|
||||
pub(crate) struct CheckpointJob {
|
||||
pub kind: CheckpointKind,
|
||||
pub presentation: PresentationDelta,
|
||||
pub context_tokens: Option<u64>,
|
||||
pub ready: Option<oneshot::Sender<std::result::Result<(), String>>>,
|
||||
}
|
||||
|
||||
pub(crate) enum CheckpointKind {
|
||||
Settled(RevisionId),
|
||||
ToolStarted {
|
||||
round_id: ToolRoundId,
|
||||
stable_revision_id: RevisionId,
|
||||
},
|
||||
ToolSettled(RevisionId),
|
||||
Final {
|
||||
revision_id: RevisionId,
|
||||
result: oneshot::Sender<Result<FinalCheckpoints>>,
|
||||
},
|
||||
Compaction {
|
||||
revision_id: RevisionId,
|
||||
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: CursorSessionHandle,
|
||||
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(revision_id)
|
||||
| CheckpointKind::ToolSettled(revision_id) => {
|
||||
publish_settled(
|
||||
&store,
|
||||
&mut builder,
|
||||
&handle,
|
||||
mode,
|
||||
revision_id,
|
||||
&presentation,
|
||||
)
|
||||
.await
|
||||
}
|
||||
CheckpointKind::ToolStarted {
|
||||
round_id,
|
||||
stable_revision_id,
|
||||
} => {
|
||||
publish_started(
|
||||
&store,
|
||||
&mut builder,
|
||||
&handle,
|
||||
mode,
|
||||
round_id,
|
||||
stable_revision_id,
|
||||
&presentation,
|
||||
)
|
||||
.await
|
||||
}
|
||||
CheckpointKind::Final {
|
||||
revision_id,
|
||||
result,
|
||||
} => {
|
||||
let checkpoints =
|
||||
build_final(&store, &mut builder, mode, revision_id, &presentation)
|
||||
.await;
|
||||
let _ = result.send(checkpoints);
|
||||
Ok(())
|
||||
}
|
||||
CheckpointKind::Compaction {
|
||||
revision_id,
|
||||
summary,
|
||||
result,
|
||||
} => {
|
||||
let messages = store.load_revision_messages(revision_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: &CursorSessionHandle,
|
||||
mode: i32,
|
||||
revision_id: RevisionId,
|
||||
presentation: &PresentationDelta,
|
||||
) -> Result<()> {
|
||||
let messages = store.load_revision_messages(revision_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: &CursorSessionHandle,
|
||||
mode: i32,
|
||||
round_id: ToolRoundId,
|
||||
stable_revision_id: RevisionId,
|
||||
presentation: &PresentationDelta,
|
||||
) -> 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_revision_messages(stable_revision_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,
|
||||
revision_id: RevisionId,
|
||||
presentation: &PresentationDelta,
|
||||
) -> Result<FinalCheckpoints> {
|
||||
let messages = store.load_revision_messages(revision_id).await?;
|
||||
let (assistant, stable) = messages
|
||||
.split_last()
|
||||
.ok_or_else(|| Error::Store("final revision 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, &PresentationDelta::default())
|
||||
.await?;
|
||||
Ok(FinalCheckpoints { staged, settled })
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
use crate::cursor::proto::agent::v1 as pb;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum CursorCommand {
|
||||
Append {
|
||||
seqno: i64,
|
||||
message: Box<pb::AgentClientMessage>,
|
||||
},
|
||||
Abort,
|
||||
Finished,
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use prost::Message;
|
||||
use tokio::sync::{oneshot, Mutex};
|
||||
|
||||
use crate::{
|
||||
cursor::{proto::agent::v1 as pb, CursorSessionHandle},
|
||||
store::{BlobId, Store},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
type ContextSender = oneshot::Sender<Result<pb::RequestContext>>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct RequestContextSynchronizer {
|
||||
handle: CursorSessionHandle,
|
||||
store: Store,
|
||||
pending: Arc<Mutex<Option<ContextSender>>>,
|
||||
}
|
||||
|
||||
impl RequestContextSynchronizer {
|
||||
pub(crate) fn new(handle: CursorSessionHandle, 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.cancellation();
|
||||
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(15)) => 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(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,389 @@
|
||||
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::{
|
||||
cursor::{
|
||||
account, analytics, bidi_append, connect, model_catalog,
|
||||
observability::CursorTraceRecorder,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
proxy::{self, CursorProxy},
|
||||
run_sse,
|
||||
},
|
||||
cursor::{CursorParent, CursorSessionRegistry},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub fn router(registry: CursorSessionRegistry) -> Result<Router> {
|
||||
let proxy = CursorProxy::cursor(registry.store().clone())?;
|
||||
Ok(router_with_proxy(registry, proxy))
|
||||
}
|
||||
|
||||
fn router_with_proxy(registry: CursorSessionRegistry, 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_append_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))
|
||||
.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<CursorSessionRegistry>,
|
||||
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 {
|
||||
super::sessions::CursorRoute::Local => {
|
||||
run_sse::stream(®istry, &request.request_id).await
|
||||
}
|
||||
super::sessions::CursorRoute::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_append_handler(
|
||||
State(registry): State<CursorSessionRegistry>,
|
||||
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_append::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().provider_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_append_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::append(®istry, 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<CursorParent>> {
|
||||
let run_id = header_text(headers, "x-parent-request-id")?;
|
||||
let tool_call_id = header_text(headers, "x-parent-agent-tool-call-id")?;
|
||||
match (run_id, tool_call_id) {
|
||||
(None, None) => Ok(None),
|
||||
(Some(run_id), Some(tool_call_id)) => Ok(Some(CursorParent {
|
||||
run_id: run_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}")))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::{body::to_bytes, routing::post};
|
||||
use prost::Message;
|
||||
use tower::ServiceExt;
|
||||
|
||||
use crate::{
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
model::ModelInvocation,
|
||||
provider::{Provider, ProviderStream},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
struct NeverProvider;
|
||||
|
||||
impl Provider for NeverProvider {
|
||||
fn stream(
|
||||
&self,
|
||||
_invocation: ModelInvocation,
|
||||
_cancellation: tokio_util::sync::CancellationToken,
|
||||
) -> ProviderStream {
|
||||
panic!("official models must not enter the BYOK provider")
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subagent_parent_headers_are_an_atomic_pair() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
"x-parent-request-id",
|
||||
HeaderValue::from_static("parent-run"),
|
||||
);
|
||||
assert!(parent_headers(&headers).is_err());
|
||||
|
||||
headers.insert(
|
||||
"x-parent-agent-tool-call-id",
|
||||
HeaderValue::from_static("parent-call"),
|
||||
);
|
||||
assert_eq!(
|
||||
parent_headers(&headers).unwrap(),
|
||||
Some(CursorParent {
|
||||
run_id: "parent-run".into(),
|
||||
tool_call_id: "parent-call".into(),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn official_model_run_sse_and_bidi_are_forwarded_together() {
|
||||
let upstream = Router::new()
|
||||
.route(
|
||||
"/agent.v1.AgentService/RunSSE",
|
||||
post(|| async { "official-stream" }),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.BidiService/BidiAppend",
|
||||
post(|| async { StatusCode::OK }),
|
||||
);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() });
|
||||
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
store.set_detailed_logging(true).await.unwrap();
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(NeverProvider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let proxy = CursorProxy::for_upstream(&format!("http://{address}")).unwrap();
|
||||
let app = router_with_proxy(registry, proxy);
|
||||
|
||||
let run = tokio::spawn(
|
||||
app.clone().oneshot(
|
||||
Request::post("/agent.v1.AgentService/RunSSE")
|
||||
.body(Body::from(
|
||||
agent::BidiRequestId {
|
||||
request_id: "official-request".into(),
|
||||
}
|
||||
.encode_to_vec(),
|
||||
))
|
||||
.unwrap(),
|
||||
),
|
||||
);
|
||||
let client_message = agent::AgentClientMessage {
|
||||
message: Some(agent::agent_client_message::Message::RunRequest(
|
||||
agent::AgentRunRequest {
|
||||
requested_model: Some(agent::RequestedModel {
|
||||
model_id: "grok-4.6".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
};
|
||||
let bidi = ai::BidiAppendRequest {
|
||||
data: hex::encode(client_message.encode_to_vec()),
|
||||
request_id: Some(ai::BidiRequestId {
|
||||
request_id: "official-request".into(),
|
||||
}),
|
||||
append_seqno: 1,
|
||||
data_binary: Vec::new(),
|
||||
};
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::post("/aiserver.v1.BidiService/BidiAppend")
|
||||
.body(Body::from(bidi.encode_to_vec()))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
|
||||
let response = run.await.unwrap().unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
to_bytes(response.into_body(), usize::MAX).await.unwrap(),
|
||||
"official-stream"
|
||||
);
|
||||
let trace = tokio::time::timeout(std::time::Duration::from_secs(2), async {
|
||||
loop {
|
||||
if let Some(trace) = store.cursor_trace("official-request").await.unwrap() {
|
||||
if trace.status == "completed" {
|
||||
break trace;
|
||||
}
|
||||
}
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(trace.route, "cursor_official");
|
||||
let artifacts = store
|
||||
.cursor_trace_artifacts("official-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let kinds = artifacts
|
||||
.iter()
|
||||
.map(|artifact| artifact.artifact_type.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert!(kinds.contains(&"bidi_append_request"));
|
||||
assert!(kinds.contains(&"run_sse_request"));
|
||||
assert!(kinds.contains(&"run_sse_chunk"));
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct OrderedInbox<T> {
|
||||
next: i64,
|
||||
pending: BTreeMap<i64, T>,
|
||||
}
|
||||
|
||||
impl<T> Default for OrderedInbox<T> {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
next: 0,
|
||||
pending: BTreeMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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 {
|
||||
return Vec::new();
|
||||
}
|
||||
self.pending.entry(seqno).or_insert(value);
|
||||
let mut ready = Vec::new();
|
||||
while let Some(value) = self.pending.remove(&self.next) {
|
||||
ready.push((self.next, value));
|
||||
self.next += 1;
|
||||
}
|
||||
ready
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
mod query;
|
||||
mod render;
|
||||
|
||||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::{ToolCall, Usage},
|
||||
provider::ModelEvent,
|
||||
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, render_tool_call, tool_completed,
|
||||
tool_placeholder, tool_started,
|
||||
};
|
||||
|
||||
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) => dynamic_mcp_placeholder(definition, call_id),
|
||||
None => 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 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(),
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
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_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),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{cursor::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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn create_plan_update_may_omit_name() {
|
||||
let message = tool_query(
|
||||
7,
|
||||
&ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-call-1".into(),
|
||||
name: "CreatePlan".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"plan": "Updated plan",
|
||||
"overview": "Update the existing plan",
|
||||
"todos": []
|
||||
}),
|
||||
},
|
||||
)
|
||||
.expect("CreatePlan updates do not require a name");
|
||||
|
||||
let Some(pb::agent_server_message::Message::InteractionQuery(query)) = message.message
|
||||
else {
|
||||
panic!("expected interaction query");
|
||||
};
|
||||
let Some(pb::interaction_query::Query::CreatePlanRequestQuery(query)) = query.query else {
|
||||
panic!("expected CreatePlan request query");
|
||||
};
|
||||
let args = query.args.expect("CreatePlan args");
|
||||
assert_eq!(args.name, "");
|
||||
assert_eq!(args.plan, "Updated plan");
|
||||
assert_eq!(args.overview, "Update the existing plan");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,575 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
proto::agent::v1 as pb,
|
||||
tools::{
|
||||
codec, edit,
|
||||
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())
|
||||
}
|
||||
"awaitshell" => Tool::AwaitToolCall(pb::AwaitToolCall::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::AwaitToolCall(tool)) => {
|
||||
tool.args = Some(pb::AwaitArgs {
|
||||
task_id: string("shell_id"),
|
||||
block_until_ms: call
|
||||
.arguments
|
||||
.get("block_until_ms")
|
||||
.and_then(Value::as_u64)
|
||||
.map(|v| v as u32),
|
||||
regex: optional("pattern"),
|
||||
})
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn direct_semble_start_uses_an_mcp_card_without_the_mcp_wrapper_shape() {
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: "SembleSearch".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"description": "Find request tracing",
|
||||
"repo": "/tmp/repo",
|
||||
"query": "request tracing"
|
||||
}),
|
||||
};
|
||||
let rendered = render_tool_call(&call, false).unwrap();
|
||||
let pb::tool_call::Tool::McpToolCall(tool) = rendered.tool.unwrap() else {
|
||||
panic!("expected MCP tool card");
|
||||
};
|
||||
let args = tool.args.unwrap();
|
||||
assert_eq!(args.server_identifier, "builtin-semble");
|
||||
assert_eq!(args.tool_name, "search");
|
||||
assert_eq!(args.name, "search");
|
||||
assert!(args.args.contains_key("repo"));
|
||||
assert!(args.args.contains_key("query"));
|
||||
assert!(!args.args.contains_key("arguments"));
|
||||
assert!(!args.args.contains_key("description"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
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("ed)
|
||||
.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())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn streams_top_level_strings_and_decodes_split_escapes() {
|
||||
let mut fields = JsonStringFields::default();
|
||||
let mut events = fields
|
||||
.push("{\"path\":\"/tmp/a\",\"count\":1,\"contents\":\"a\\n\\uD8")
|
||||
.unwrap();
|
||||
events.extend(fields.push("3D\\uDE00b\"}").unwrap());
|
||||
assert_eq!(
|
||||
events,
|
||||
vec![
|
||||
StringFieldEvent::Delta {
|
||||
name: "path".into(),
|
||||
text: "/tmp/a".into()
|
||||
},
|
||||
StringFieldEvent::End {
|
||||
name: "path".into()
|
||||
},
|
||||
StringFieldEvent::Delta {
|
||||
name: "contents".into(),
|
||||
text: "a\n".into()
|
||||
},
|
||||
StringFieldEvent::Delta {
|
||||
name: "contents".into(),
|
||||
text: "😀b".into()
|
||||
},
|
||||
StringFieldEvent::End {
|
||||
name: "contents".into()
|
||||
},
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::CursorSessionHandle,
|
||||
cursor::{
|
||||
connect::{
|
||||
encode_end_stream, encode_error_end_stream, ConnectCode, ConnectErrorDetail,
|
||||
ConnectStreamError,
|
||||
},
|
||||
proto::aiserver::v1 as ai,
|
||||
},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
pub fn finish_success(handle: &CursorSessionHandle) {
|
||||
handle.emit_frame(encode_end_stream());
|
||||
handle.close_output();
|
||||
}
|
||||
|
||||
pub fn fail(handle: &CursorSessionHandle, error: &Error) -> Result<()> {
|
||||
let stream_error = match error {
|
||||
Error::Provider(_) | Error::Http(_) => provider_error(error),
|
||||
Error::Protocol(message) => plain_message(ConnectCode::InvalidArgument, message.clone()),
|
||||
Error::Decode(_) | Error::Json(_) => plain_error(ConnectCode::InvalidArgument, error),
|
||||
Error::RunNotFound(_) => plain_error(ConnectCode::NotFound, error),
|
||||
Error::Cancelled => plain_error(ConnectCode::Canceled, error),
|
||||
Error::Config(_)
|
||||
| Error::Store(_)
|
||||
| Error::Database(_)
|
||||
| Error::Migration(_)
|
||||
| Error::Encode(_)
|
||||
| Error::Io(_) => plain_error(ConnectCode::Internal, error),
|
||||
};
|
||||
handle.emit_frame(encode_error_end_stream(&stream_error)?);
|
||||
handle.close_output();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn cancel(handle: &CursorSessionHandle) -> Result<()> {
|
||||
handle.emit_frame(encode_error_end_stream(&ConnectStreamError {
|
||||
code: ConnectCode::Canceled,
|
||||
message: "run was cancelled".into(),
|
||||
details: Vec::new(),
|
||||
})?);
|
||||
handle.close_output();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn plain_error(code: ConnectCode, error: &Error) -> ConnectStreamError {
|
||||
plain_message(code, error.to_string())
|
||||
}
|
||||
|
||||
fn plain_message(code: ConnectCode, message: String) -> ConnectStreamError {
|
||||
ConnectStreamError {
|
||||
code,
|
||||
message,
|
||||
details: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_error(error: &Error) -> ConnectStreamError {
|
||||
let detail = ai::ErrorDetails {
|
||||
error: ai::error_details::Error::ProviderError as i32,
|
||||
details: Some(ai::CustomErrorDetails {
|
||||
title: "Provider Error".into(),
|
||||
detail: error.to_string(),
|
||||
allow_command_links_potentially_unsafe_please_only_use_for_handwritten_trusted_markdown:
|
||||
Some(true),
|
||||
is_retryable: Some(true),
|
||||
show_request_id: Some(true),
|
||||
should_show_immediate_error: Some(false),
|
||||
}),
|
||||
is_expected: Some(true),
|
||||
};
|
||||
ConnectStreamError {
|
||||
code: ConnectCode::Unavailable,
|
||||
message: error.to_string(),
|
||||
details: vec![ConnectErrorDetail {
|
||||
type_name: "aiserver.v1.ErrorDetails".into(),
|
||||
value: STANDARD_NO_PAD.encode(detail.encode_to_vec()),
|
||||
}],
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
mod account;
|
||||
mod actor;
|
||||
mod analytics;
|
||||
pub mod bidi_append;
|
||||
pub mod blob_sync;
|
||||
pub mod checkpoint;
|
||||
pub mod connect;
|
||||
mod context_sync;
|
||||
pub mod handlers;
|
||||
mod inbox;
|
||||
pub mod interaction;
|
||||
mod json_stream;
|
||||
pub(crate) mod lifecycle;
|
||||
mod model_catalog;
|
||||
pub(crate) mod observability;
|
||||
mod presentation;
|
||||
mod projection;
|
||||
pub mod prompting;
|
||||
pub mod proto;
|
||||
pub mod proxy;
|
||||
pub mod request;
|
||||
pub mod run_sse;
|
||||
pub mod session;
|
||||
pub mod sessions;
|
||||
pub mod tools;
|
||||
mod usage;
|
||||
|
||||
pub use command::CursorCommand;
|
||||
pub use sessions::{CursorParent, CursorSessionHandle, CursorSessionRegistry};
|
||||
mod command;
|
||||
@@ -0,0 +1,701 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
extract::{Extension, State},
|
||||
http::{header, HeaderValue, Request, Response, StatusCode},
|
||||
};
|
||||
use bytes::{BufMut, BytesMut};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
proto::agent::v1 as agent,
|
||||
proxy::{self, CursorProxy},
|
||||
CursorSessionRegistry,
|
||||
},
|
||||
model::ProviderModel,
|
||||
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";
|
||||
|
||||
pub async fn available_models(
|
||||
State(registry): State<CursorSessionRegistry>,
|
||||
Extension(proxy): Extension<CursorProxy>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let models = registry.store().provider_models(true).await?;
|
||||
let provider_names = registry
|
||||
.store()
|
||||
.providers()
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|provider| (provider.provider_id, provider.name))
|
||||
.collect::<HashMap<_, _>>();
|
||||
tracing::info!(
|
||||
model_count = models.len(),
|
||||
"appending BYOK models to Cursor AvailableModels"
|
||||
);
|
||||
let available_models = models
|
||||
.iter()
|
||||
.map(|model| {
|
||||
let provider_name = provider_names.get(&model.provider_id).ok_or_else(|| {
|
||||
Error::Config(format!(
|
||||
"provider {} for model {} does not exist",
|
||||
model.provider_id, model.model_hash
|
||||
))
|
||||
})?;
|
||||
Ok(available_model(model, provider_name))
|
||||
})
|
||||
.collect::<Result<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<CursorSessionRegistry>,
|
||||
Extension(proxy): Extension<CursorProxy>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let models = registry.store().provider_models(true).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: &ProviderModel, provider_name: &str) -> AvailableModel {
|
||||
let variants = model_variants(model);
|
||||
let legacy_slugs = variants
|
||||
.iter()
|
||||
.filter_map(|variant| variant.legacy_slug.clone())
|
||||
.collect();
|
||||
let tooltip = model_tooltip(model, "200K", "high", false);
|
||||
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(),
|
||||
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: provider_name.into(),
|
||||
variant: 1,
|
||||
dismiss_on_selection: false,
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
fn model_parameters() -> 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
|
||||
.into_iter()
|
||||
.map(|(value, display_name)| EnumParameterValue {
|
||||
value: value.into(),
|
||||
display_name: Some(display_name.into()),
|
||||
})
|
||||
.collect(),
|
||||
}),
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(false),
|
||||
},
|
||||
ModelParameterDefinition {
|
||||
id: "effort".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: &ProviderModel) -> 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: &ProviderModel,
|
||||
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: "effort".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, context_name, effort, fast)),
|
||||
display_name_outside_picker: Some(display_name),
|
||||
variant_string_representation: Some(format!(
|
||||
"{}[context={context},effort={effort},fast={fast}]",
|
||||
model.model_hash
|
||||
)),
|
||||
legacy_slug: Some(format!(
|
||||
"{}-{context}-{effort}{}",
|
||||
model.model_hash,
|
||||
if fast { "-fast" } else { "" }
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn model_tooltip(
|
||||
model: &ProviderModel,
|
||||
context_name: &str,
|
||||
effort: &str,
|
||||
fast: bool,
|
||||
) -> TooltipData {
|
||||
let fast_label = if fast { " (Fast)" } else { "" };
|
||||
TooltipData {
|
||||
markdown_content: Some(format!(
|
||||
"**{}{fast_label}**<br /><br />{context_name} context window<br /><br />*Version: {effort} effort*",
|
||||
model.display_name
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn usable_model(model: &ProviderModel) -> 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()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::body::{to_bytes, Bytes};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn maps_byok_model_to_cursor_catalog_fields() {
|
||||
let model = ProviderModel {
|
||||
model_hash: "33ceed20".into(),
|
||||
provider_id: 1,
|
||||
model_id: "deepseek-v4-flash".into(),
|
||||
display_name: "DeepSeek V4 Flash".into(),
|
||||
endpoint_type: crate::model::ProviderType::OpenAiResponses,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: Some(200_000),
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
created_at_ms: 0,
|
||||
updated_at_ms: 0,
|
||||
};
|
||||
|
||||
let mapped = available_model(&model, "OpenRouter");
|
||||
assert_eq!(mapped.name, "33ceed20");
|
||||
assert!(mapped.default_on);
|
||||
assert_eq!(mapped.supports_agent, Some(true));
|
||||
assert_eq!(mapped.degradation_status, Some(0));
|
||||
assert_eq!(mapped.supports_thinking, Some(true));
|
||||
assert_eq!(mapped.supports_images, Some(true));
|
||||
assert_eq!(mapped.supports_max_mode, Some(true));
|
||||
assert_eq!(mapped.supports_non_max_mode, Some(true));
|
||||
assert_eq!(mapped.supports_plan_mode, Some(true));
|
||||
assert_eq!(mapped.supports_sandboxing, Some(true));
|
||||
assert_eq!(mapped.supports_cmd_k, Some(false));
|
||||
assert_eq!(
|
||||
mapped.client_display_name.as_deref(),
|
||||
Some("DeepSeek V4 Flash")
|
||||
);
|
||||
assert_eq!(mapped.server_model_name.as_deref(), Some("33ceed20"));
|
||||
assert_eq!(mapped.named_model_section_index, Some(1));
|
||||
assert_eq!(mapped.vendor_name.as_deref(), Some("cursor"));
|
||||
assert_eq!(mapped.parameter_definitions.len(), 3);
|
||||
let context = mapped
|
||||
.parameter_definitions
|
||||
.iter()
|
||||
.find(|parameter| parameter.id == "context")
|
||||
.unwrap();
|
||||
let context_values = context
|
||||
.parameter_type
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.enum_parameter
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.values
|
||||
.iter()
|
||||
.map(|value| value.value.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(context_values, ["200k", "356k", "800k", "1m"]);
|
||||
let effort = mapped
|
||||
.parameter_definitions
|
||||
.iter()
|
||||
.find(|parameter| parameter.id == "effort")
|
||||
.unwrap();
|
||||
assert!(effort
|
||||
.parameter_type
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.enum_parameter
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.values
|
||||
.iter()
|
||||
.any(|value| value.value == "max"));
|
||||
assert_eq!(mapped.variants.len(), 40);
|
||||
assert_eq!(mapped.legacy_slugs.len(), 40);
|
||||
assert_eq!(mapped.model_picker_badges.len(), 1);
|
||||
assert_eq!(mapped.model_picker_badges[0].label, "OpenRouter");
|
||||
assert!(!mapped.model_picker_badges[0].dismiss_on_selection);
|
||||
let default = mapped
|
||||
.variants
|
||||
.iter()
|
||||
.find(|variant| variant.is_default_non_max_config == Some(true))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
default.variant_string_representation.as_deref(),
|
||||
Some("33ceed20[context=200k,effort=high,fast=false]")
|
||||
);
|
||||
assert_eq!(mapped.vendor.unwrap().display_name, "Cursor");
|
||||
assert!(usable_model(&model).thinking_details.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn appends_models_without_reencoding_official_fields() {
|
||||
// Unknown field 99 = 7 stands in for every official field this service does not know.
|
||||
let official = Bytes::from_static(&[0x98, 0x06, 0x07]);
|
||||
let addition = AvailableModelsAddition {
|
||||
model_names: vec!["f246010a".into()],
|
||||
models: Vec::new(),
|
||||
}
|
||||
.encode_to_vec();
|
||||
let response = merge_response(
|
||||
proxy::BufferedResponse {
|
||||
status: axum::http::StatusCode::OK,
|
||||
headers: Default::default(),
|
||||
body: official.clone(),
|
||||
},
|
||||
addition.clone(),
|
||||
)
|
||||
.unwrap();
|
||||
let merged = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
assert_eq!(&merged[..official.len()], official.as_ref());
|
||||
assert_eq!(&merged[official.len()..], addition);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn updates_connect_length_when_catalog_is_framed() {
|
||||
let official = [0x98, 0x06, 0x07];
|
||||
let mut framed = BytesMut::new();
|
||||
framed.put_u8(0);
|
||||
framed.put_u32(official.len() as u32);
|
||||
framed.extend_from_slice(&official);
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert(axum::http::header::CONTENT_LENGTH, framed.len().into());
|
||||
let response = merge_response(
|
||||
proxy::BufferedResponse {
|
||||
status: axum::http::StatusCode::OK,
|
||||
headers,
|
||||
body: framed.freeze(),
|
||||
},
|
||||
vec![0x0a, 0x01, b'x'],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(response.headers()[axum::http::header::CONTENT_LENGTH], "11");
|
||||
let merged = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
assert_eq!(u32::from_be_bytes(merged[1..5].try_into().unwrap()), 6);
|
||||
assert_eq!(&merged[5..8], &official);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn returns_local_catalog_when_upstream_rejects_request() {
|
||||
let local = AvailableModelsAddition {
|
||||
model_names: vec!["f246010a".into()],
|
||||
models: Vec::new(),
|
||||
}
|
||||
.encode_to_vec();
|
||||
let response = merge_response(
|
||||
proxy::BufferedResponse {
|
||||
status: axum::http::StatusCode::UNAUTHORIZED,
|
||||
headers: Default::default(),
|
||||
body: Bytes::from_static(b"not logged in"),
|
||||
},
|
||||
local.clone(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), axum::http::StatusCode::OK);
|
||||
assert_eq!(
|
||||
response.headers()[axum::http::header::CONTENT_TYPE],
|
||||
"application/proto"
|
||||
);
|
||||
assert_eq!(
|
||||
to_bytes(response.into_body(), usize::MAX).await.unwrap(),
|
||||
local
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
use crate::{store::BlobId, store::Store};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorTraceRecorder {
|
||||
store: Store,
|
||||
request_id: String,
|
||||
}
|
||||
|
||||
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(),
|
||||
}),
|
||||
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(),
|
||||
}),
|
||||
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]) {
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.add_cursor_trace_response_chunk(&self.request_id, source, data)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn finish(&self, error: Option<&str>) {
|
||||
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");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::cursor::{proto::agent::v1 as pb, tools::result::ToolCompletion};
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct PresentationDelta {
|
||||
pub steps: Vec<pb::ConversationStep>,
|
||||
pub read_paths: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct Presentation {
|
||||
steps: Vec<pb::ConversationStep>,
|
||||
read_paths: Vec<String>,
|
||||
text: String,
|
||||
thinking: String,
|
||||
}
|
||||
|
||||
impl Presentation {
|
||||
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 take(&mut self) -> PresentationDelta {
|
||||
PresentationDelta {
|
||||
steps: std::mem::take(&mut self.steps),
|
||||
read_paths: std::mem::take(&mut self.read_paths),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn thinking_step_keeps_the_measured_duration() {
|
||||
let mut presentation = Presentation::default();
|
||||
presentation.thinking_delta("reasoning");
|
||||
presentation.finish_thinking(Duration::from_millis(6_880));
|
||||
let step = presentation.take().steps.pop().unwrap();
|
||||
let Some(pb::conversation_step::Message::ThinkingMessage(thinking)) = step.message else {
|
||||
panic!("expected thinking step");
|
||||
};
|
||||
assert_eq!(thinking.text, "reasoning");
|
||||
assert_eq!(thinking.duration_ms, 6_880);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
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 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 wire_id.starts_with("request-context:")
|
||||
|| wire_id.starts_with("selected-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() {
|
||||
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,276 @@
|
||||
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 {
|
||||
"AwaitShell" => "AWAIT".into(),
|
||||
"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",
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn direct_semble_tools_use_the_cursor_mcp_execution_contract() {
|
||||
let dynamic = HashSet::new();
|
||||
assert_eq!(tool_identifier("SembleSearch", &dynamic), "MCP");
|
||||
assert_eq!(tool_identifier("SembleFindRelated", &dynamic), "MCP");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
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,243 @@
|
||||
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 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);
|
||||
}
|
||||
@@ -0,0 +1,207 @@
|
||||
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] = &[
|
||||
"REQUEST_CONTEXT",
|
||||
"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")
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn direct_semble_tools_are_available_in_every_working_mode() {
|
||||
let assets = PromptAssets::embedded().unwrap();
|
||||
for mode in [
|
||||
Mode::Agent,
|
||||
Mode::Ask,
|
||||
Mode::Plan,
|
||||
Mode::Debug,
|
||||
Mode::Multitask,
|
||||
Mode::Subagent,
|
||||
] {
|
||||
let names = assets
|
||||
.mode(mode)
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert!(names.contains(&"SembleSearch"), "missing in {mode:?}");
|
||||
assert!(names.contains(&"SembleFindRelated"), "missing in {mode:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
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(super) struct Catalog {
|
||||
tools: HashMap<String, ToolDefinition>,
|
||||
variants: HashMap<String, ToolDefinition>,
|
||||
}
|
||||
|
||||
impl Catalog {
|
||||
pub(super) 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(super) fn select_json(&self, manifest: &str) -> Result<Vec<ToolDefinition>> {
|
||||
let manifest: Manifest = serde_json::from_str(manifest)?;
|
||||
self.select(&manifest)
|
||||
}
|
||||
|
||||
fn select(&self, manifest: &Manifest) -> Result<Vec<ToolDefinition>> {
|
||||
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()))?,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
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,47 @@
|
||||
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(input),
|
||||
"createplan" | "updateplan" | "writeplan" => state.plan = Some(input),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
state
|
||||
}
|
||||
|
||||
fn normalize(value: &str) -> String {
|
||||
value
|
||||
.chars()
|
||||
.filter(|character| character.is_ascii_alphanumeric())
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
mod assets;
|
||||
mod catalog;
|
||||
mod compiler;
|
||||
mod derived_state;
|
||||
|
||||
pub use assets::*;
|
||||
pub use compiler::*;
|
||||
pub use derived_state::*;
|
||||
@@ -0,0 +1,71 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
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(),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn for_upstream(upstream: &str) -> Result<Self> {
|
||||
Ok(Self {
|
||||
client: Some(
|
||||
reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()?,
|
||||
),
|
||||
store: None,
|
||||
upstream: upstream.trim_end_matches('/').to_owned(),
|
||||
})
|
||||
}
|
||||
|
||||
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>> {
|
||||
let started = Instant::now();
|
||||
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);
|
||||
|
||||
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::harness::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");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
extract::Extension,
|
||||
http::{header, Request, StatusCode},
|
||||
response::IntoResponse,
|
||||
routing::any,
|
||||
Router,
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::{forward, CursorProxy};
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_request_and_response() {
|
||||
let upstream = Router::new().route(
|
||||
"/unknown",
|
||||
any(|request: Request<Body>| async move {
|
||||
let method = request.method().clone();
|
||||
let query = request.uri().query().unwrap_or_default().to_owned();
|
||||
let marker = request.headers()["x-marker"].clone();
|
||||
let body = to_bytes(request.into_body(), usize::MAX).await.unwrap();
|
||||
(
|
||||
StatusCode::CREATED,
|
||||
[(header::CONTENT_TYPE, "application/proto")],
|
||||
format!(
|
||||
"{method} {query} {} {}",
|
||||
marker.to_str().unwrap(),
|
||||
String::from_utf8_lossy(&body)
|
||||
),
|
||||
)
|
||||
.into_response()
|
||||
}),
|
||||
);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() });
|
||||
let proxy = CursorProxy::for_upstream(&format!("http://{address}")).unwrap();
|
||||
let app = Router::new().fallback(forward).layer(Extension(proxy));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::put("/unknown?a=1")
|
||||
.header("x-marker", "kept")
|
||||
.body(Body::from("payload"))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::CREATED);
|
||||
assert_eq!(
|
||||
response.headers()[header::CONTENT_TYPE],
|
||||
"application/proto"
|
||||
);
|
||||
assert_eq!(
|
||||
to_bytes(response.into_body(), usize::MAX).await.unwrap(),
|
||||
"PUT a=1 kept payload"
|
||||
);
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,330 @@
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use crate::{cursor::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 identities = BTreeSet::new();
|
||||
let mut contexts = Vec::with_capacity(action.completions.len());
|
||||
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 {
|
||||
return Err(Error::Protocol(format!(
|
||||
"background task notification is not a finished task: {}",
|
||||
reason.as_str_name()
|
||||
)));
|
||||
}
|
||||
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 identity = agent_id.unwrap_or(&completion.task_id);
|
||||
let identity = format!("{}:{identity}", kind.as_str_name());
|
||||
if !identities.insert(identity.clone()) {
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate background task completion: {identity}"
|
||||
)));
|
||||
}
|
||||
contexts.push(completion_context(completion, kind, agent_id)?);
|
||||
}
|
||||
|
||||
let first = &action.completions[0];
|
||||
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: contexts.join("\n\n"),
|
||||
turn_user: pb::UserMessage {
|
||||
text,
|
||||
message_id: format!(
|
||||
"background-completed:{}",
|
||||
identities.into_iter().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!(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn finished_subagent_becomes_an_idempotent_user_runtime_event() {
|
||||
let action = pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![completion()],
|
||||
};
|
||||
let projection = project(&action, pb::AgentMode::Multitask as i32).unwrap();
|
||||
|
||||
assert!(projection.context.contains("kind: subagent"));
|
||||
assert!(projection.context.contains("agent_id: child-id"));
|
||||
assert!(projection.context.contains("child result"));
|
||||
|
||||
assert_eq!(projection.turn_user.text, FOLLOW_UP);
|
||||
assert_eq!(projection.turn_user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
projection.turn_user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn finished_shell_becomes_the_captured_system_notification() {
|
||||
let action = pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![shell_completion()],
|
||||
};
|
||||
let projection = project(&action, pb::AgentMode::Agent as i32).unwrap();
|
||||
|
||||
assert_eq!(projection.turn_user.text, SHELL_FOLLOW_UP);
|
||||
assert_eq!(
|
||||
projection.context,
|
||||
concat!(
|
||||
"<system_notification>\n",
|
||||
"The following task has finished. If you were already aware, ignore this notification and do not restate prior responses.\n\n",
|
||||
"<task>\n",
|
||||
"kind: shell\n",
|
||||
"status: aborted\n",
|
||||
"task_id: 977679\n",
|
||||
"title: Start Python HTTP server on 9000\n",
|
||||
"tool_call_id: shell-call\n",
|
||||
"detail: terminated_by_user\n",
|
||||
"output_path: /tmp/977679.txt\n",
|
||||
"thread_id: terminal-thread\n",
|
||||
"</task>\n",
|
||||
"</system_notification>"
|
||||
)
|
||||
);
|
||||
assert_eq!(projection.turn_user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
projection.turn_user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
let metadata = projection.turn_user.simulated_message_metadata.unwrap();
|
||||
assert_eq!(
|
||||
metadata.title.as_deref(),
|
||||
Some("Start Python HTTP server on 9000")
|
||||
);
|
||||
assert_eq!(metadata.task_id.as_deref(), Some("977679"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shell_and_subagent_completions_keep_both_follow_up_contracts() {
|
||||
let projection = project(
|
||||
&pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![shell_completion(), completion()],
|
||||
},
|
||||
pb::AgentMode::Multitask as i32,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(projection.context.contains("kind: shell"));
|
||||
assert!(projection.context.contains("agent_id: child-id"));
|
||||
assert!(projection.turn_user.text.contains(SHELL_FOLLOW_UP));
|
||||
assert!(projection.turn_user.text.contains(FOLLOW_UP));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_requires_the_captured_subagent_identity_and_terminal_reason() {
|
||||
let mut value = completion();
|
||||
value.subagent_id = None;
|
||||
assert!(project(
|
||||
&pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![value]
|
||||
},
|
||||
pb::AgentMode::Agent as i32
|
||||
)
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("subagent_id"));
|
||||
|
||||
let mut value = completion();
|
||||
value.reason = pb::BackgroundTaskCompletionReason::TaskProgress as i32;
|
||||
assert!(project(
|
||||
&pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![value]
|
||||
},
|
||||
pb::AgentMode::Agent as i32
|
||||
)
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("not a finished task"));
|
||||
}
|
||||
|
||||
fn completion() -> pb::BackgroundTaskCompletion {
|
||||
pb::BackgroundTaskCompletion {
|
||||
task_id: "child-id".into(),
|
||||
kind: pb::BackgroundTaskKind::Subagent as i32,
|
||||
status: pb::BackgroundTaskStatus::Success as i32,
|
||||
title: "Inspect protocol".into(),
|
||||
detail: Some("child result".into()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
subagent_id: Some("child-id".into()),
|
||||
tool_call_id: Some("task-call".into()),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
fn shell_completion() -> pb::BackgroundTaskCompletion {
|
||||
pb::BackgroundTaskCompletion {
|
||||
task_id: "977679".into(),
|
||||
kind: pb::BackgroundTaskKind::Shell as i32,
|
||||
status: pb::BackgroundTaskStatus::Aborted as i32,
|
||||
title: "Start Python HTTP server on 9000".into(),
|
||||
detail: Some("terminated_by_user".into()),
|
||||
output_path: Some("/tmp/977679.txt".into()),
|
||||
thread_id: Some("terminal-thread".into()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
tool_call_id: Some("shell-call".into()),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,542 @@
|
||||
use std::{
|
||||
collections::{BTreeMap, HashMap, HashSet},
|
||||
path::Path,
|
||||
};
|
||||
|
||||
use prost::Message;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
context_sync::RequestContextSynchronizer, proto::agent::v1 as pb, 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 definition = ToolDefinition {
|
||||
name: wire.name.clone(),
|
||||
description: wire.description.clone(),
|
||||
parameters,
|
||||
};
|
||||
if output
|
||||
.insert(wire.name.clone(), (wire.clone(), definition))
|
||||
.is_some()
|
||||
{
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate MCP tool definition: {}",
|
||||
wire.name
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
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('&', "&")
|
||||
.replace('"', """)
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn meta_mcp_routes_projects_descriptor_routing_without_runtime_discovery() {
|
||||
let context = pb::RequestContext {
|
||||
mcp_meta_tool_options: Some(pb::McpMetaToolOptions {
|
||||
enabled: true,
|
||||
mcp_descriptors: vec![pb::McpDescriptor {
|
||||
server_name: "fast-context".into(),
|
||||
server_identifier: "fast-context".into(),
|
||||
tools: vec![pb::McpToolDescriptor {
|
||||
tool_name: "fast_context_search".into(),
|
||||
description: Some("search code".into()),
|
||||
input_schema_json: Some(r#"{"type":"object"}"#.into()),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
}],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let routes = meta_mcp_routes(&context);
|
||||
let route = routes
|
||||
.get(&("fast-context".into(), "fast_context_search".into()))
|
||||
.unwrap();
|
||||
assert_eq!(route.name, "fast-context-fast_context_search");
|
||||
assert_eq!(route.provider_identifier, "fast-context");
|
||||
assert_eq!(route.tool_name, "fast_context_search");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
use crate::{
|
||||
cursor::{blob_sync::BlobSynchronizer, proto::agent::v1 as pb},
|
||||
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,9 @@
|
||||
mod background;
|
||||
mod context;
|
||||
mod images;
|
||||
mod model;
|
||||
mod prepare;
|
||||
mod runtime;
|
||||
|
||||
pub use prepare::*;
|
||||
pub(crate) use runtime::compile_injection;
|
||||
@@ -0,0 +1,276 @@
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::{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(¶meter.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
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_token_count(value: &str) -> Option<u64> {
|
||||
let value = value.trim().to_ascii_lowercase();
|
||||
let (number, multiplier) = match value.chars().last()? {
|
||||
'k' => (&value[..value.len() - 1], 1_000),
|
||||
'm' => (&value[..value.len() - 1], 1_000_000),
|
||||
_ => (value.as_str(), 1),
|
||||
};
|
||||
number.parse::<u64>().ok()?.checked_mul(multiplier)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn requested(id: &str, parameters: &[(&str, &str)]) -> pb::RequestedModel {
|
||||
pb::RequestedModel {
|
||||
model_id: id.into(),
|
||||
parameters: parameters
|
||||
.iter()
|
||||
.map(|(id, value)| pb::requested_model::ModelParameterValue {
|
||||
id: (*id).into(),
|
||||
value: (*value).into(),
|
||||
})
|
||||
.collect(),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cursor_model_parameters_keep_order_and_define_reasoning() {
|
||||
let model = from_requested(
|
||||
&requested("grok-4.6", &[("effort", "xhigh"), ("fast", "false")]),
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(model.model_id, "grok-4.6");
|
||||
assert!(model.reasoning.enabled);
|
||||
assert_eq!(model.reasoning.effort.as_deref(), Some("xhigh"));
|
||||
assert_eq!(model.latency, ModelLatency::Standard);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cursor_reasoning_and_context_metadata_are_normalized() {
|
||||
let model = from_requested(
|
||||
&requested(
|
||||
"gpt-5.6-sol",
|
||||
&[("context", "272k"), ("reasoning", "medium")],
|
||||
),
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(model.context_window_tokens, Some(272_000));
|
||||
assert_eq!(model.reasoning.effort.as_deref(), Some("medium"));
|
||||
assert!(model.reasoning.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cursor_catalog_context_values_are_consumed() {
|
||||
for (value, tokens) in [
|
||||
("200k", 200_000),
|
||||
("356k", 356_000),
|
||||
("800k", 800_000),
|
||||
("1m", 1_000_000),
|
||||
] {
|
||||
let model = from_requested(&requested("model", &[("context", value)]), None).unwrap();
|
||||
assert_eq!(model.context_window_tokens, Some(tokens));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subagent_override_distinguishes_explicit_inherit_and_disabled() {
|
||||
let request = pb::AgentRunRequest {
|
||||
subagent_model_overrides: vec![
|
||||
pb::SubagentModelOverride {
|
||||
subagent_type: "explore".into(),
|
||||
selection: Some(pb::subagent_model_override::Selection::Model(requested(
|
||||
"claude-opus-5",
|
||||
&[("thinking", "true")],
|
||||
))),
|
||||
},
|
||||
pb::SubagentModelOverride {
|
||||
subagent_type: "generalPurpose".into(),
|
||||
selection: Some(pb::subagent_model_override::Selection::Inherit(true)),
|
||||
},
|
||||
pb::SubagentModelOverride {
|
||||
subagent_type: "shell".into(),
|
||||
selection: Some(pb::subagent_model_override::Selection::Disabled(true)),
|
||||
},
|
||||
],
|
||||
..Default::default()
|
||||
};
|
||||
let overrides = overrides(&request).unwrap();
|
||||
assert!(matches!(
|
||||
&overrides[0],
|
||||
(SubagentKind::Named(name), SubagentModelOverride::Explicit(model))
|
||||
if name == "explore" && model.reasoning.enabled
|
||||
));
|
||||
assert!(matches!(
|
||||
&overrides[1],
|
||||
(SubagentKind::GeneralPurpose, SubagentModelOverride::Inherit)
|
||||
));
|
||||
assert!(matches!(
|
||||
&overrides[2],
|
||||
(SubagentKind::Named(name), SubagentModelOverride::Disabled) if name == "shell"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_subagent_model_is_inherit() {
|
||||
let request = pb::AgentRunRequest {
|
||||
subagent_model_overrides: vec![pb::SubagentModelOverride {
|
||||
subagent_type: "generalPurpose".into(),
|
||||
selection: Some(pb::subagent_model_override::Selection::Model(requested(
|
||||
"default",
|
||||
&[],
|
||||
))),
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
assert!(matches!(
|
||||
overrides(&request).unwrap().as_slice(),
|
||||
[(SubagentKind::GeneralPurpose, SubagentModelOverride::Inherit)]
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cursor_only_parameters_do_not_leak_into_model_spec() {
|
||||
let model = from_requested(
|
||||
&requested("grok-4.6", &[("fast", "true"), ("context", "300k")]),
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(model.latency, ModelLatency::Fast);
|
||||
assert_eq!(model.context_window_tokens, Some(300_000));
|
||||
assert!(from_requested(&requested("grok-4.6", &[("mystery", "x")]), None).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,639 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::{
|
||||
cursor::prompting::{Mode, PromptCompiler},
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer,
|
||||
checkpoint::CheckpointBuilder,
|
||||
context_sync::RequestContextSynchronizer,
|
||||
projection,
|
||||
proto::agent::v1 as pb,
|
||||
tools::runtime::{ExecContext, SubagentModel},
|
||||
},
|
||||
model::{
|
||||
CanonicalMessage, ContentPart, ConversationId, MessageContent, Origin, PreparedRun,
|
||||
PromptSpec, Role, RunAction, RunId, RunKind,
|
||||
},
|
||||
store::{BlobId, Store},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{background, context, model, runtime};
|
||||
|
||||
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(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,
|
||||
parent: Option<(RunId, String)>,
|
||||
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()),
|
||||
);
|
||||
// RunSSE/Bidi request_id identifies this concrete execution attempt. Cursor may
|
||||
// reuse AgentRunRequest.run_id when a queued or subagent-driven attempt resumes.
|
||||
let run_id = RunId::new(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,
|
||||
event_id,
|
||||
input_id,
|
||||
starts_turn,
|
||||
compacting,
|
||||
background_completion,
|
||||
} = action(request_id, 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(provider_model) = store
|
||||
.provider_model(&model.model_id)
|
||||
.await?
|
||||
.filter(|model| model.enabled)
|
||||
{
|
||||
provider_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_revision_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_revision(&conversation_id, messages).await?
|
||||
}
|
||||
Some(_) | None => store.ensure_conversation(&conversation_id).await?,
|
||||
};
|
||||
let base_revision_id = match input_id {
|
||||
Some(input_id) => {
|
||||
store
|
||||
.anchor_input(&conversation_id, &input_id, proposed_base_revision_id)
|
||||
.await?
|
||||
}
|
||||
None => proposed_base_revision_id,
|
||||
};
|
||||
let existing_runtime = match event_id.as_deref() {
|
||||
Some(event_id) if !background_completion => {
|
||||
store
|
||||
.message(&conversation_id, &format!("runtime:{event_id}"))
|
||||
.await?
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let 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) = runtime::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)) => match existing_runtime {
|
||||
Some(message) => vec![message],
|
||||
None => vec![
|
||||
runtime::compile(
|
||||
event_id,
|
||||
checkpoint_mode,
|
||||
&user,
|
||||
&request_context,
|
||||
&action_context,
|
||||
compiler,
|
||||
blob_sync,
|
||||
)
|
||||
.await?,
|
||||
],
|
||||
},
|
||||
(None, None) => Vec::new(),
|
||||
_ => {
|
||||
return Err(Error::Protocol(
|
||||
"Cursor action has an incomplete runtime event".into(),
|
||||
))
|
||||
}
|
||||
}
|
||||
};
|
||||
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(projection::decode_pending(pending)?),
|
||||
pending => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"Cursor resume contains {} pending assistant messages",
|
||||
pending.len()
|
||||
)))
|
||||
}
|
||||
};
|
||||
RunAction::Resume { pending_tool_round }
|
||||
};
|
||||
let kind = match (request.subagent_type_name.as_deref(), parent) {
|
||||
(None, _) => RunKind::Root,
|
||||
(Some(name), Some((parent_run_id, parent_tool_call_id))) => RunKind::Subagent {
|
||||
parent_run_id,
|
||||
parent_tool_call_id,
|
||||
kind: model::subagent_kind(name),
|
||||
background: false,
|
||||
},
|
||||
(Some(_), None) => {
|
||||
return Err(Error::Protocol(
|
||||
"subagent Run is missing its parent Run and tool call".into(),
|
||||
));
|
||||
}
|
||||
};
|
||||
let exec = exec_context(
|
||||
request,
|
||||
&request_context,
|
||||
&conversation_id,
|
||||
&model.model_id,
|
||||
subagents_disabled,
|
||||
&subagent_model_overrides,
|
||||
);
|
||||
Ok((
|
||||
PreparedRun {
|
||||
run_id,
|
||||
conversation_id,
|
||||
kind,
|
||||
model,
|
||||
prompt,
|
||||
initial_messages,
|
||||
action,
|
||||
base_revision_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,
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
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 action(request_id: &str, request: &pb::AgentRunRequest) -> Result<ActionProjection> {
|
||||
let mode = request
|
||||
.conversation_state
|
||||
.as_ref()
|
||||
.and_then(|state| state.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())
|
||||
})?;
|
||||
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: user.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(),
|
||||
);
|
||||
Ok(ActionProjection {
|
||||
mode: user.mode,
|
||||
turn_user: Some(user.clone()),
|
||||
action_context: context.join("\n\n"),
|
||||
event_id: Some(format!("run-request:{request_id}")),
|
||||
input_id: Some(format!("cursor:user:{}", user.message_id)),
|
||||
starts_turn: true,
|
||||
compacting: false,
|
||||
background_completion: false,
|
||||
})
|
||||
}
|
||||
pb::conversation_action::Action::BackgroundTaskCompletionAction(action) => {
|
||||
let projection = background::project(action, mode)?;
|
||||
Ok(ActionProjection {
|
||||
mode,
|
||||
action_context: projection.context,
|
||||
event_id: Some(format!("run-request:{request_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),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn restored_system_root_is_structural_not_bound_to_the_next_model() {
|
||||
let prompt = CanonicalMessage::text(
|
||||
"root",
|
||||
Role::System,
|
||||
Origin::Prompt,
|
||||
"prompt from the previous model",
|
||||
);
|
||||
validate_prompt_root(std::slice::from_ref(&prompt)).unwrap();
|
||||
assert!(validate_prompt_root(&[prompt.clone(), prompt]).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unsupported_cursor_mode_is_not_silently_treated_as_agent() {
|
||||
assert_eq!(
|
||||
mode_from_proto(pb::AgentMode::Agent as i32).unwrap(),
|
||||
Mode::Agent
|
||||
);
|
||||
assert!(mode_from_proto(pb::AgentMode::Project as i32).is_err());
|
||||
assert!(mode_from_proto(99).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn current_user_message_consumes_the_mode_instead_of_history_mode() {
|
||||
let request = pb::AgentRunRequest {
|
||||
conversation_state: Some(pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "explain".into(),
|
||||
message_id: "user-message".into(),
|
||||
mode: pb::AgentMode::Ask as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
let projection = action("request", &request).unwrap();
|
||||
assert_eq!(projection.mode, pb::AgentMode::Ask as i32);
|
||||
assert_eq!(
|
||||
projection.input_id.as_deref(),
|
||||
Some("cursor:user:user-message")
|
||||
);
|
||||
assert_eq!(mode_from_proto(projection.mode).unwrap(), Mode::Ask);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execute_plan_appends_the_approved_plan_as_a_stable_runtime_event() {
|
||||
let execute = pb::ExecutePlanAction {
|
||||
plan_file_uri: Some("file:///workspace/example.plan.md".into()),
|
||||
plan_file_content: Some("# Build\n\n- implement it".into()),
|
||||
execution_mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
let request = pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::ExecutePlanAction(
|
||||
execute.clone(),
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let first = action("request-one", &request).unwrap();
|
||||
let second = action("request-two", &request).unwrap();
|
||||
assert_eq!(first.mode, pb::AgentMode::Agent as i32);
|
||||
assert!(first.starts_turn);
|
||||
assert_eq!(first.event_id, second.event_id);
|
||||
assert_eq!(first.input_id, None);
|
||||
assert_eq!(
|
||||
first.turn_user.as_ref().map(|user| user.text.as_str()),
|
||||
Some("Execute the approved plan.")
|
||||
);
|
||||
assert!(first
|
||||
.action_context
|
||||
.contains("file:///workspace/example.plan.md"));
|
||||
assert!(first.action_context.contains("# Build\n\n- implement it"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execute_plan_requires_content() {
|
||||
let result = execute_plan(&pb::ExecutePlanAction {
|
||||
execution_mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
});
|
||||
assert!(matches!(
|
||||
result,
|
||||
Err(Error::Protocol(message)) if message.contains("missing plan content")
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use chrono::{Offset, Utc};
|
||||
use chrono_tz::Tz;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer,
|
||||
prompting::{Mode, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
},
|
||||
model::{CanonicalMessage, MessageContent, Origin, Role},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{context, images};
|
||||
|
||||
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::prepare::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 time = Time::now(
|
||||
request_context
|
||||
.env
|
||||
.as_ref()
|
||||
.map(|env| env.time_zone.as_str()),
|
||||
)?;
|
||||
let mut values = BTreeMap::from([
|
||||
(
|
||||
"REQUEST_CONTEXT",
|
||||
section(context::compile_context(request_context, &time.today)),
|
||||
),
|
||||
("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", time.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 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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
use axum::{
|
||||
body::Body,
|
||||
http::{header, HeaderValue, Response, StatusCode},
|
||||
};
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::StreamExt;
|
||||
|
||||
use crate::{
|
||||
cursor::{observability::CursorTraceRecorder, CursorSessionRegistry},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub async fn stream(registry: &CursorSessionRegistry, request_id: &str) -> Result<Response<Body>> {
|
||||
let handle = registry.get_or_create(request_id).await?;
|
||||
let mut receiver = handle.subscribe();
|
||||
let trace = handle.trace().cloned();
|
||||
if let Some(trace) = &trace {
|
||||
trace.response_started(StatusCode::OK.as_u16()).await;
|
||||
}
|
||||
let body_stream = async_stream::stream! {
|
||||
let mut trace = TraceStreamSink::new(trace, "byok_server");
|
||||
while let Some(chunk) = receiver.recv().await {
|
||||
trace.chunk(&chunk);
|
||||
yield Ok::<Bytes, std::convert::Infallible>(chunk);
|
||||
}
|
||||
trace.finish(None);
|
||||
};
|
||||
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)
|
||||
}
|
||||
|
||||
pub async fn upstream(
|
||||
registry: CursorSessionRegistry,
|
||||
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) {
|
||||
self.finish(None);
|
||||
}
|
||||
}
|
||||
|
||||
struct UpstreamRunGuard {
|
||||
registry: CursorSessionRegistry,
|
||||
request_id: String,
|
||||
generation: u64,
|
||||
}
|
||||
|
||||
impl Drop for UpstreamRunGuard {
|
||||
fn drop(&mut self) {
|
||||
self.registry
|
||||
.finish_upstream(self.request_id.clone(), self.generation);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,698 @@
|
||||
use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
|
||||
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::{
|
||||
client::{ClientCommand, ClientEvent, ClientSession, CommitCause},
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer,
|
||||
checkpoint::{
|
||||
worker::{CheckpointJob, CheckpointKind, CheckpointWorker, FinalCheckpoints},
|
||||
CheckpointBuilder,
|
||||
},
|
||||
interaction,
|
||||
presentation::Presentation,
|
||||
prompting::PromptCompiler,
|
||||
proto::agent::v1 as pb,
|
||||
request::CursorRunContext,
|
||||
tools::{
|
||||
codec,
|
||||
result::{ToolCompletion, ToolResultReceiver},
|
||||
runtime::CursorToolRuntime,
|
||||
stream::ToolCallStream,
|
||||
ToolBatchState, ToolDispatcher,
|
||||
},
|
||||
},
|
||||
model::{ToolCall, ToolRoundId, Usage},
|
||||
run::{RunFailure, RunOutcome},
|
||||
store::{Store, ToolRoundStatus},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::CursorSessionHandle;
|
||||
|
||||
pub struct CursorSession {
|
||||
handle: CursorSessionHandle,
|
||||
store: Store,
|
||||
context: CursorRunContext,
|
||||
core: ClientSession,
|
||||
tools: ToolDispatcher,
|
||||
results: ToolResultReceiver,
|
||||
checkpoint: CheckpointBuilder,
|
||||
tool_runtime: CursorToolRuntime,
|
||||
runtime_actions: mpsc::UnboundedReceiver<pb::InjectContextAction>,
|
||||
compiler: PromptCompiler,
|
||||
blob_sync: BlobSynchronizer,
|
||||
injection_ids: HashSet<String>,
|
||||
pending_injections: HashMap<String, PendingInjection>,
|
||||
}
|
||||
|
||||
struct PendingInjection {
|
||||
user_message: Option<pb::UserMessage>,
|
||||
delivery_batch_id: String,
|
||||
}
|
||||
|
||||
pub(crate) struct CursorSessionRuntime {
|
||||
pub tools: ToolDispatcher,
|
||||
pub results: ToolResultReceiver,
|
||||
pub checkpoint: CheckpointBuilder,
|
||||
pub tool_runtime: CursorToolRuntime,
|
||||
pub runtime_actions: mpsc::UnboundedReceiver<pb::InjectContextAction>,
|
||||
pub compiler: PromptCompiler,
|
||||
pub blob_sync: BlobSynchronizer,
|
||||
}
|
||||
|
||||
impl CursorSession {
|
||||
pub(crate) fn new(
|
||||
handle: CursorSessionHandle,
|
||||
store: Store,
|
||||
context: CursorRunContext,
|
||||
core: ClientSession,
|
||||
runtime: CursorSessionRuntime,
|
||||
) -> Self {
|
||||
Self {
|
||||
handle,
|
||||
store,
|
||||
context,
|
||||
core,
|
||||
tools: runtime.tools,
|
||||
results: runtime.results,
|
||||
checkpoint: runtime.checkpoint,
|
||||
tool_runtime: runtime.tool_runtime,
|
||||
runtime_actions: runtime.runtime_actions,
|
||||
compiler: runtime.compiler,
|
||||
blob_sync: runtime.blob_sync,
|
||||
injection_ids: HashSet::new(),
|
||||
pending_injections: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(mut self) -> Result<()> {
|
||||
if self.context.compacting {
|
||||
self.handle.emit(&interaction::summary_started())?;
|
||||
}
|
||||
let mut worker = CheckpointWorker::spawn(
|
||||
self.store.clone(),
|
||||
self.checkpoint.clone(),
|
||||
self.handle.clone(),
|
||||
self.context.mode,
|
||||
);
|
||||
let mut checkpoint_worker_open = true;
|
||||
let mut calls = BTreeMap::<usize, ToolCall>::new();
|
||||
let mut streams = BTreeMap::<usize, ToolCallStream>::new();
|
||||
let mut completions = HashMap::<String, ToolCompletion>::new();
|
||||
let mut completed = HashSet::<String>::new();
|
||||
let mut response_text = String::new();
|
||||
let mut response_thinking = String::new();
|
||||
let mut active_round = None::<ToolRoundId>;
|
||||
let mut final_checkpoint = None::<FinalCheckpoints>;
|
||||
let mut compaction_checkpoint = None::<pb::ConversationStateStructure>;
|
||||
let mut turn_usage = None::<Usage>;
|
||||
let mut context_tokens = None::<u64>;
|
||||
let mut ready = VecDeque::new();
|
||||
let mut presentation = Presentation::default();
|
||||
|
||||
loop {
|
||||
let input = if let Some(completion) = ready.pop_front() {
|
||||
Input::Completion(completion)
|
||||
} else {
|
||||
tokio::select! {
|
||||
event = self.core.events.recv() => Input::Event(event),
|
||||
completion = self.results.recv() => Input::CompletionResult(completion),
|
||||
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
||||
failure = worker.failures.recv(), if checkpoint_worker_open => Input::CheckpointFailure(failure),
|
||||
}
|
||||
};
|
||||
match input {
|
||||
Input::CheckpointFailure(Some(error)) => return Err(error),
|
||||
Input::CheckpointFailure(None) => {
|
||||
checkpoint_worker_open = false;
|
||||
}
|
||||
Input::Completion(completion) => {
|
||||
self.forward_completion(completion, &mut completions)
|
||||
.await?;
|
||||
}
|
||||
Input::CompletionResult(Some(result)) => {
|
||||
self.forward_completion(result?, &mut completions).await?;
|
||||
}
|
||||
Input::CompletionResult(None) => {
|
||||
return Err(Error::Protocol("tool result channel closed".into()));
|
||||
}
|
||||
Input::RuntimeAction(Some(action)) => {
|
||||
self.forward_injection(*action).await?;
|
||||
}
|
||||
Input::RuntimeAction(None) => {
|
||||
return Err(Error::Protocol("runtime action channel closed".into()));
|
||||
}
|
||||
Input::Event(None) => {
|
||||
worker.abort();
|
||||
return Err(Error::Protocol("core event channel closed".into()));
|
||||
}
|
||||
Input::Event(Some(event)) => match event {
|
||||
ClientEvent::AutoCompactionStarted => {
|
||||
self.handle.emit(&interaction::summary_started())?;
|
||||
}
|
||||
ClientEvent::AutoCompactionCompleted => {
|
||||
self.handle.emit(&interaction::summary_completed())?;
|
||||
}
|
||||
ClientEvent::TextStart => {}
|
||||
ClientEvent::TextEnd => {
|
||||
if !self.context.compacting {
|
||||
presentation.finish_text();
|
||||
}
|
||||
}
|
||||
ClientEvent::TextDelta(delta) => {
|
||||
response_text.push_str(&delta);
|
||||
if self.context.compacting {
|
||||
self.handle.emit(&interaction::summary_delta(delta))?;
|
||||
} else {
|
||||
presentation.text_delta(&delta);
|
||||
self.emit_model_event(
|
||||
crate::provider::ModelEvent::TextDelta(delta),
|
||||
"",
|
||||
)?;
|
||||
}
|
||||
}
|
||||
ClientEvent::ThinkingStart => {}
|
||||
ClientEvent::ThinkingDelta(delta) => {
|
||||
response_thinking.push_str(&delta);
|
||||
if !self.context.compacting {
|
||||
presentation.thinking_delta(&delta);
|
||||
self.emit_model_event(
|
||||
crate::provider::ModelEvent::ThinkingDelta(delta),
|
||||
"",
|
||||
)?;
|
||||
}
|
||||
}
|
||||
ClientEvent::ThinkingEnd { duration } => {
|
||||
if !self.context.compacting {
|
||||
presentation.finish_thinking(duration);
|
||||
self.handle
|
||||
.emit(&interaction::thinking_completed(duration))?;
|
||||
}
|
||||
}
|
||||
ClientEvent::ToolCallStart {
|
||||
index,
|
||||
call_id,
|
||||
name,
|
||||
model_call_id,
|
||||
} => {
|
||||
let call = ToolCall {
|
||||
index,
|
||||
call_id: call_id.clone(),
|
||||
model_call_id: model_call_id.clone(),
|
||||
name: name.clone(),
|
||||
arguments_text: String::new(),
|
||||
arguments: serde_json::Value::Null,
|
||||
};
|
||||
self.emit_model_event(
|
||||
crate::provider::ModelEvent::ToolCallStart {
|
||||
index,
|
||||
call_id,
|
||||
name: name.clone(),
|
||||
},
|
||||
&model_call_id,
|
||||
)?;
|
||||
streams.insert(
|
||||
index,
|
||||
ToolCallStream::new(&name, self.context.dynamic_tools.get(&name)),
|
||||
);
|
||||
calls.insert(index, call);
|
||||
}
|
||||
ClientEvent::ToolCallArgumentsDelta { index, delta } => {
|
||||
let call = calls.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("unknown streaming tool index: {index}"))
|
||||
})?;
|
||||
call.arguments_text.push_str(&delta);
|
||||
let stream = streams.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
||||
})?;
|
||||
for message in stream.arguments_delta(call, &delta)? {
|
||||
self.handle.emit(&message)?;
|
||||
}
|
||||
}
|
||||
ClientEvent::ToolCallEnd { index } => {
|
||||
let call = calls.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("unknown completed tool index: {index}"))
|
||||
})?;
|
||||
call.arguments = serde_json::from_str(&call.arguments_text)?;
|
||||
}
|
||||
ClientEvent::Usage(usage) => {
|
||||
if !self.context.compacting {
|
||||
if let Some(output_tokens) = usage.output_tokens {
|
||||
self.handle.emit(&interaction::token_delta(output_tokens))?;
|
||||
}
|
||||
}
|
||||
if !self.context.compacting {
|
||||
context_tokens = usage
|
||||
.input_tokens
|
||||
.zip(usage.output_tokens)
|
||||
.and_then(|(input, output)| input.checked_add(output));
|
||||
}
|
||||
match &mut turn_usage {
|
||||
Some(total) => *total += usage,
|
||||
None => turn_usage = Some(usage),
|
||||
}
|
||||
}
|
||||
ClientEvent::ExecuteToolRound {
|
||||
round_id,
|
||||
calls: round_calls,
|
||||
} => {
|
||||
active_round = Some(round_id);
|
||||
for dispatched in self
|
||||
.tools
|
||||
.start_batch(
|
||||
&round_calls,
|
||||
ToolBatchState {
|
||||
completed: &completed,
|
||||
started: &HashSet::new(),
|
||||
response_text: &response_text,
|
||||
response_thinking: &response_thinking,
|
||||
},
|
||||
&self
|
||||
.store
|
||||
.load_current_messages(&crate::model::ConversationId::new(
|
||||
&self.context.exec.conversation_id,
|
||||
))
|
||||
.await?,
|
||||
&self.context.dynamic_tools,
|
||||
&self.context.exec,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
for message in dispatched.messages {
|
||||
self.handle.emit(&message)?;
|
||||
}
|
||||
if let Some(completion) = dispatched.completion {
|
||||
ready.push_back(completion);
|
||||
}
|
||||
}
|
||||
response_text.clear();
|
||||
response_thinking.clear();
|
||||
calls.clear();
|
||||
streams.clear();
|
||||
}
|
||||
ClientEvent::StateCommitted(state) => {
|
||||
if matches!(&state.cause, CommitCause::RuntimeEvent { .. }) {
|
||||
response_text.clear();
|
||||
response_thinking.clear();
|
||||
calls.clear();
|
||||
streams.clear();
|
||||
}
|
||||
if let CommitCause::RuntimeEvent { event_id } = &state.cause {
|
||||
if let Some(injection_id) = event_id.strip_prefix("inject-context:") {
|
||||
if let Some(pending) = self.pending_injections.remove(injection_id)
|
||||
{
|
||||
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
|
||||
.min(i64::MAX as u64)
|
||||
as i64;
|
||||
self.handle.emit(&interaction::context_injection_delivered(
|
||||
injection_id.to_owned(),
|
||||
pending.delivery_batch_id.clone(),
|
||||
delivered_at_ms,
|
||||
))?;
|
||||
if let Some(user_message) = pending.user_message {
|
||||
self.handle.emit(&interaction::user_message_appended(
|
||||
user_message,
|
||||
))?;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if let CommitCause::ToolRoundStarted(round_id) = &state.cause {
|
||||
active_round = Some(round_id.clone());
|
||||
}
|
||||
let mut tool_round_settled = false;
|
||||
if let CommitCause::ToolResult { call_id } = &state.cause {
|
||||
let completion = completions.remove(call_id).ok_or_else(|| {
|
||||
Error::Protocol(format!(
|
||||
"core committed a tool result without typed Cursor state: {call_id}"
|
||||
))
|
||||
})?;
|
||||
let snapshot = self
|
||||
.store
|
||||
.tool_round(active_round.as_ref().ok_or_else(|| {
|
||||
Error::Protocol("tool commit has no active round".into())
|
||||
})?)
|
||||
.await?
|
||||
.ok_or_else(|| {
|
||||
Error::Store("active tool round disappeared".into())
|
||||
})?;
|
||||
let call = snapshot
|
||||
.calls
|
||||
.iter()
|
||||
.find(|call| call.call_id == *call_id)
|
||||
.ok_or_else(|| {
|
||||
Error::Protocol(format!(
|
||||
"committed call is absent from tool round: {call_id}"
|
||||
))
|
||||
})?;
|
||||
self.handle
|
||||
.emit(&interaction::tool_completed(call, &completion))?;
|
||||
presentation.tool_completed(&completion);
|
||||
completed.insert(call_id.clone());
|
||||
tool_round_settled = snapshot.status == ToolRoundStatus::Settled;
|
||||
}
|
||||
let final_turn = state.cause == CommitCause::FinalTurn;
|
||||
if let CommitCause::Compaction { summary } = &state.cause {
|
||||
if !state.barrier.is_required() {
|
||||
return Err(Error::Protocol(
|
||||
"compaction state has no completion barrier".into(),
|
||||
));
|
||||
}
|
||||
let (sender, receiver) = oneshot::channel();
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::Compaction {
|
||||
revision_id: state.revision_id,
|
||||
summary: summary.clone(),
|
||||
result: sender,
|
||||
},
|
||||
presentation: presentation.take(),
|
||||
context_tokens: None,
|
||||
ready: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
match receiver
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||
{
|
||||
Ok(checkpoint) => {
|
||||
compaction_checkpoint = Some(checkpoint);
|
||||
state.barrier.complete(Ok(()));
|
||||
}
|
||||
Err(error) => {
|
||||
state.barrier.complete(Err(error.to_string()));
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if final_turn {
|
||||
if !state.barrier.is_required() {
|
||||
return Err(Error::Protocol(
|
||||
"final state has no completion barrier".into(),
|
||||
));
|
||||
}
|
||||
let (sender, receiver) = oneshot::channel();
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::Final {
|
||||
revision_id: state.revision_id,
|
||||
result: sender,
|
||||
},
|
||||
presentation: presentation.take(),
|
||||
context_tokens,
|
||||
ready: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
match receiver
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||
{
|
||||
Ok(checkpoints) => {
|
||||
final_checkpoint = Some(checkpoints);
|
||||
state.barrier.complete(Ok(()));
|
||||
}
|
||||
Err(error) => {
|
||||
state.barrier.complete(Err(error.to_string()));
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
} else if let CommitCause::ToolRoundStarted(round_id) = &state.cause {
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::ToolStarted {
|
||||
round_id: round_id.clone(),
|
||||
stable_revision_id: state.revision_id,
|
||||
},
|
||||
presentation: presentation.take(),
|
||||
context_tokens,
|
||||
ready: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
} else if tool_round_settled {
|
||||
if !state.barrier.is_required() {
|
||||
return Err(Error::Protocol(
|
||||
"settled tool round has no completion barrier".into(),
|
||||
));
|
||||
}
|
||||
let (ready, published) = oneshot::channel();
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::ToolSettled(state.revision_id),
|
||||
presentation: presentation.take(),
|
||||
context_tokens,
|
||||
ready: Some(ready),
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
let result = published
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||
.map_err(Error::Protocol);
|
||||
match result {
|
||||
Ok(()) => state.barrier.complete(Ok(())),
|
||||
Err(error) => {
|
||||
state.barrier.complete(Err(error.to_string()));
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
active_round = None;
|
||||
self.tool_runtime.clear_completed().await;
|
||||
} else if !matches!(&state.cause, CommitCause::ToolResult { .. })
|
||||
&& active_round.is_some()
|
||||
{
|
||||
let round_id = active_round.clone().ok_or_else(|| {
|
||||
Error::Protocol("active tool round disappeared".into())
|
||||
})?;
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::ToolStarted {
|
||||
round_id,
|
||||
stable_revision_id: state.revision_id,
|
||||
},
|
||||
presentation: presentation.take(),
|
||||
context_tokens,
|
||||
ready: None,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
} else if !matches!(&state.cause, CommitCause::ToolResult { .. }) {
|
||||
let requires_ready = state.barrier.is_required();
|
||||
let (ready, published) = oneshot::channel();
|
||||
worker
|
||||
.jobs
|
||||
.send(CheckpointJob {
|
||||
kind: CheckpointKind::Settled(state.revision_id),
|
||||
presentation: presentation.take(),
|
||||
context_tokens,
|
||||
ready: requires_ready.then_some(ready),
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Protocol("checkpoint worker closed".into()))?;
|
||||
if requires_ready {
|
||||
let result = published
|
||||
.await
|
||||
.map_err(|_| {
|
||||
Error::Protocol("checkpoint worker stopped".into())
|
||||
})?
|
||||
.map_err(Error::Protocol);
|
||||
match result {
|
||||
Ok(()) => state.barrier.complete(Ok(())),
|
||||
Err(error) => {
|
||||
state.barrier.complete(Err(error.to_string()));
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
ClientEvent::Ended(outcome) => {
|
||||
return match outcome {
|
||||
RunOutcome::Completed => {
|
||||
if self.context.compacting {
|
||||
let checkpoint =
|
||||
compaction_checkpoint.take().ok_or_else(|| {
|
||||
Error::Protocol(
|
||||
"Completed compaction without checkpoint".into(),
|
||||
)
|
||||
})?;
|
||||
self.handle.emit(&interaction::summary_completed())?;
|
||||
self.handle.emit(&interaction::turn_ended(turn_usage))?;
|
||||
for _ in 0..3 {
|
||||
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
||||
}
|
||||
crate::cursor::lifecycle::finish_success(&self.handle);
|
||||
return Ok(());
|
||||
}
|
||||
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
||||
Error::Protocol("Completed without final state".into())
|
||||
})?;
|
||||
self.handle.emit(&interaction::turn_ended(turn_usage))?;
|
||||
self.checkpoint
|
||||
.publish(&self.handle, &checkpoints.staged)
|
||||
.await?;
|
||||
self.checkpoint
|
||||
.publish(&self.handle, &checkpoints.settled)
|
||||
.await?;
|
||||
self.handle.emit(&pb::AgentServerMessage {
|
||||
ttft_breakdown: None,
|
||||
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
||||
})?;
|
||||
crate::cursor::lifecycle::finish_success(&self.handle);
|
||||
Ok(())
|
||||
}
|
||||
RunOutcome::Cancelled => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
crate::cursor::lifecycle::cancel(&self.handle)
|
||||
}
|
||||
RunOutcome::Failed(failure) => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
crate::cursor::lifecycle::fail(&self.handle, &cursor_error(failure))
|
||||
}
|
||||
};
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn abort_execs(&self) {
|
||||
for id in self.tool_runtime.drain_running().await {
|
||||
let _ = self.handle.emit(&codec::abort(id));
|
||||
}
|
||||
}
|
||||
|
||||
async fn forward_completion(
|
||||
&self,
|
||||
mut completion: ToolCompletion,
|
||||
completions: &mut HashMap<String, ToolCompletion>,
|
||||
) -> Result<()> {
|
||||
if let Some(image) = completion.take_read_image() {
|
||||
let blob_id = self.store.put_blob(&image.data, &[]).await?;
|
||||
completion.persist_read_image(&blob_id, &image)?;
|
||||
}
|
||||
let result = completion.result();
|
||||
if result.call_id.is_empty() {
|
||||
return Err(Error::Protocol("tool result call_id is empty".into()));
|
||||
}
|
||||
if completions
|
||||
.insert(result.call_id.clone(), completion.clone())
|
||||
.is_some()
|
||||
{
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate tool result call_id: {}",
|
||||
result.call_id
|
||||
)));
|
||||
}
|
||||
self.core
|
||||
.commands
|
||||
.send(ClientCommand::ToolResult(result.clone()))
|
||||
.await
|
||||
.map_err(|_| Error::RunNotFound(self.context.request_id.clone()))
|
||||
}
|
||||
|
||||
async fn forward_injection(&mut self, action: pb::InjectContextAction) -> Result<()> {
|
||||
if action.injection_id.is_empty() {
|
||||
return Err(Error::Protocol(
|
||||
"InjectContextAction has no injection_id".into(),
|
||||
));
|
||||
}
|
||||
if action.expected_run_id != self.context.request_id {
|
||||
return Err(Error::Protocol(format!(
|
||||
"InjectContextAction expected run {}, active run is {}",
|
||||
action.expected_run_id, self.context.request_id
|
||||
)));
|
||||
}
|
||||
if self.injection_ids.contains(&action.injection_id) {
|
||||
return Ok(());
|
||||
}
|
||||
let user_message = match action.payload.as_ref() {
|
||||
Some(pb::inject_context_action::Payload::UserContext(context)) => {
|
||||
context.user_message.clone()
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let message = crate::cursor::request::compile_injection(
|
||||
&action,
|
||||
self.context.mode,
|
||||
&self.compiler,
|
||||
&self.blob_sync,
|
||||
)
|
||||
.await?;
|
||||
let injection_id = action.injection_id;
|
||||
let delivery_batch_id = injection_id.clone();
|
||||
self.injection_ids.insert(injection_id.clone());
|
||||
self.pending_injections.insert(
|
||||
injection_id.clone(),
|
||||
PendingInjection {
|
||||
user_message,
|
||||
delivery_batch_id,
|
||||
},
|
||||
);
|
||||
self.handle
|
||||
.emit(&interaction::context_injection_queued(injection_id.clone()))?;
|
||||
if self
|
||||
.core
|
||||
.commands
|
||||
.send(ClientCommand::RuntimeMessage(message))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
self.pending_injections.remove(&injection_id);
|
||||
return Err(Error::RunNotFound(self.context.request_id.clone()));
|
||||
}
|
||||
self.interrupt_execs().await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn interrupt_execs(&self) {
|
||||
// Keep runtime entries until Cursor returns the aborted result. The core tool
|
||||
// round needs that terminal result before it can append the injected context
|
||||
// after the complete assistant/tool pair and continue the same Run.
|
||||
for id in self.tool_runtime.running_exec_ids().await {
|
||||
let _ = self.handle.emit(&codec::abort(id));
|
||||
}
|
||||
}
|
||||
|
||||
fn emit_model_event(
|
||||
&self,
|
||||
event: crate::provider::ModelEvent,
|
||||
model_call_id: &str,
|
||||
) -> Result<()> {
|
||||
if let Some(message) =
|
||||
interaction::response_event(&event, model_call_id, &self.context.dynamic_tools)?
|
||||
{
|
||||
self.handle.emit(&message)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
enum Input {
|
||||
Event(Option<ClientEvent>),
|
||||
Completion(ToolCompletion),
|
||||
CompletionResult(Option<Result<ToolCompletion>>),
|
||||
RuntimeAction(Option<Box<pb::InjectContextAction>>),
|
||||
CheckpointFailure(Option<Error>),
|
||||
}
|
||||
|
||||
fn cursor_error(failure: RunFailure) -> Error {
|
||||
match failure {
|
||||
RunFailure::Protocol(message) => Error::Protocol(message),
|
||||
RunFailure::Provider(message) => Error::Provider(message),
|
||||
RunFailure::Store(message) => Error::Store(message),
|
||||
RunFailure::Client(message) => Error::Protocol(message),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,307 @@
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::{Arc, OnceLock},
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::{mpsc, Mutex, Notify};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
cursor::prompting::PromptCompiler,
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer, observability::CursorTraceRecorder, proto::agent::v1 as pb,
|
||||
},
|
||||
provider::Provider,
|
||||
run::RunRegistry,
|
||||
store::Store,
|
||||
Result,
|
||||
};
|
||||
|
||||
use super::{
|
||||
actor::{CursorActor, RunDependencies},
|
||||
CursorCommand,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorSessionHandle {
|
||||
request_id: String,
|
||||
commands: mpsc::Sender<CursorCommand>,
|
||||
output: Arc<OutputHub>,
|
||||
cancellation: CancellationToken,
|
||||
parent: Arc<OnceLock<CursorParent>>,
|
||||
trace: Option<CursorTraceRecorder>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct CursorParent {
|
||||
pub run_id: String,
|
||||
pub tool_call_id: String,
|
||||
}
|
||||
|
||||
impl CursorSessionHandle {
|
||||
pub fn request_id(&self) -> &str {
|
||||
&self.request_id
|
||||
}
|
||||
pub fn subscribe(&self) -> mpsc::UnboundedReceiver<Bytes> {
|
||||
self.output.subscribe()
|
||||
}
|
||||
pub async fn command(&self, command: CursorCommand) -> Result<()> {
|
||||
self.commands
|
||||
.send(command)
|
||||
.await
|
||||
.map_err(|_| crate::Error::RunNotFound(self.request_id.clone()))
|
||||
}
|
||||
pub fn emit_frame(&self, frame: Bytes) {
|
||||
self.output.emit(frame);
|
||||
}
|
||||
pub fn emit(&self, message: &pb::AgentServerMessage) -> Result<()> {
|
||||
self.emit_frame(crate::cursor::connect::encode_message(message)?);
|
||||
Ok(())
|
||||
}
|
||||
pub fn cancel(&self) {
|
||||
self.cancellation.cancel();
|
||||
}
|
||||
pub fn close_output(&self) {
|
||||
self.output.close();
|
||||
}
|
||||
pub fn cancellation(&self) -> CancellationToken {
|
||||
self.cancellation.clone()
|
||||
}
|
||||
pub fn set_parent(&self, parent: CursorParent) -> Result<()> {
|
||||
if parent.run_id.is_empty() || parent.tool_call_id.is_empty() {
|
||||
return Err(crate::Error::Protocol(
|
||||
"Cursor parent run and tool call ids are required".into(),
|
||||
));
|
||||
}
|
||||
if self.parent.get().is_some_and(|current| current != &parent) {
|
||||
return Err(crate::Error::Protocol(format!(
|
||||
"conflicting parent ids for request {}",
|
||||
self.request_id
|
||||
)));
|
||||
}
|
||||
let _ = self.parent.set(parent);
|
||||
Ok(())
|
||||
}
|
||||
pub fn parent(&self) -> Option<&CursorParent> {
|
||||
self.parent.get()
|
||||
}
|
||||
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
|
||||
self.trace.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct OutputHub {
|
||||
state: parking_lot::Mutex<OutputState>,
|
||||
closed: tokio::sync::Notify,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct OutputState {
|
||||
history: Vec<Bytes>,
|
||||
subscribers: Vec<mpsc::UnboundedSender<Bytes>>,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
impl OutputHub {
|
||||
fn emit(&self, frame: Bytes) {
|
||||
let mut state = self.state.lock();
|
||||
if state.closed {
|
||||
return;
|
||||
}
|
||||
state.history.push(frame.clone());
|
||||
state
|
||||
.subscribers
|
||||
.retain(|subscriber| subscriber.send(frame.clone()).is_ok());
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
fn close(&self) {
|
||||
let mut state = self.state.lock();
|
||||
state.closed = true;
|
||||
state.subscribers.clear();
|
||||
drop(state);
|
||||
self.closed.notify_waiters();
|
||||
}
|
||||
|
||||
async fn wait_closed(&self) {
|
||||
loop {
|
||||
let notified = self.closed.notified();
|
||||
if self.state.lock().closed {
|
||||
return;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorSessionRegistry {
|
||||
inner: Arc<RegistryInner>,
|
||||
}
|
||||
|
||||
struct RegistryInner {
|
||||
runs: Mutex<HashMap<String, CursorSessionHandle>>,
|
||||
upstream_runs: Mutex<HashMap<String, u64>>,
|
||||
route_changed: Notify,
|
||||
run_registry: RunRegistry,
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum CursorRoute {
|
||||
Local,
|
||||
Upstream(u64),
|
||||
}
|
||||
|
||||
impl CursorSessionRegistry {
|
||||
pub fn store(&self) -> &Store {
|
||||
&self.inner.store
|
||||
}
|
||||
|
||||
pub fn new(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
run_registry: RunRegistry,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
runs: Mutex::new(HashMap::new()),
|
||||
upstream_runs: Mutex::new(HashMap::new()),
|
||||
route_changed: Notify::new(),
|
||||
run_registry,
|
||||
store,
|
||||
provider,
|
||||
compiler,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_or_create(&self, request_id: &str) -> Result<CursorSessionHandle> {
|
||||
if let Some(handle) = self.inner.runs.lock().await.get(request_id).cloned() {
|
||||
return Ok(handle);
|
||||
}
|
||||
let (commands, receiver) = mpsc::channel(128);
|
||||
let output = Arc::new(OutputHub::default());
|
||||
let cancellation = CancellationToken::new();
|
||||
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await;
|
||||
let handle = CursorSessionHandle {
|
||||
request_id: request_id.into(),
|
||||
commands,
|
||||
output,
|
||||
cancellation,
|
||||
parent: Arc::new(OnceLock::new()),
|
||||
trace,
|
||||
};
|
||||
let mut runs = self.inner.runs.lock().await;
|
||||
if let Some(existing) = runs.get(request_id).cloned() {
|
||||
return Ok(existing);
|
||||
}
|
||||
runs.insert(request_id.into(), handle.clone());
|
||||
drop(runs);
|
||||
self.inner.route_changed.notify_waiters();
|
||||
let blob_sync =
|
||||
BlobSynchronizer::new(request_id.into(), self.inner.store.clone(), handle.clone());
|
||||
CursorActor::spawn(
|
||||
handle.clone(),
|
||||
receiver,
|
||||
RunDependencies {
|
||||
store: self.inner.store.clone(),
|
||||
provider: self.inner.provider.clone(),
|
||||
compiler: self.inner.compiler.clone(),
|
||||
run_registry: self.inner.run_registry.clone(),
|
||||
},
|
||||
blob_sync,
|
||||
0,
|
||||
);
|
||||
let registry = Arc::downgrade(&self.inner);
|
||||
let request_id = request_id.to_string();
|
||||
let output = handle.output.clone();
|
||||
tokio::spawn(async move {
|
||||
output.wait_closed().await;
|
||||
let Some(registry) = registry.upgrade() else {
|
||||
return;
|
||||
};
|
||||
registry.runs.lock().await.remove(&request_id);
|
||||
});
|
||||
Ok(handle)
|
||||
}
|
||||
|
||||
pub(crate) async fn local(&self, request_id: &str) -> Option<CursorSessionHandle> {
|
||||
self.inner.runs.lock().await.get(request_id).cloned()
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_upstream(&self, request_id: &str) {
|
||||
let mut runs = self.inner.upstream_runs.lock().await;
|
||||
let generation = runs.get(request_id).copied().unwrap_or_default() + 1;
|
||||
runs.insert(request_id.into(), generation);
|
||||
drop(runs);
|
||||
self.inner.route_changed.notify_waiters();
|
||||
}
|
||||
|
||||
pub(crate) async fn upstream(&self, request_id: &str) -> bool {
|
||||
self.inner
|
||||
.upstream_runs
|
||||
.lock()
|
||||
.await
|
||||
.contains_key(request_id)
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_route(&self, request_id: &str) -> CursorRoute {
|
||||
loop {
|
||||
let changed = self.inner.route_changed.notified();
|
||||
if self.inner.runs.lock().await.contains_key(request_id) {
|
||||
return CursorRoute::Local;
|
||||
}
|
||||
if let Some(generation) = self
|
||||
.inner
|
||||
.upstream_runs
|
||||
.lock()
|
||||
.await
|
||||
.get(request_id)
|
||||
.copied()
|
||||
{
|
||||
return CursorRoute::Upstream(generation);
|
||||
}
|
||||
changed.await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn finish_upstream(&self, request_id: String, generation: u64) {
|
||||
let registry = self.clone();
|
||||
tokio::spawn(async move {
|
||||
let mut runs = registry.inner.upstream_runs.lock().await;
|
||||
if runs.get(&request_id) == Some(&generation) {
|
||||
runs.remove(&request_id);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
pub async fn shutdown(&self) {
|
||||
let handles = {
|
||||
let mut runs = self.inner.runs.lock().await;
|
||||
runs.drain().map(|(_, handle)| handle).collect::<Vec<_>>()
|
||||
};
|
||||
self.inner.run_registry.shutdown().await;
|
||||
self.inner.upstream_runs.lock().await.clear();
|
||||
for handle in handles {
|
||||
handle.cancel();
|
||||
let _ = crate::cursor::lifecycle::cancel(&handle);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
mod request;
|
||||
mod response;
|
||||
|
||||
pub use request::{abort, mcp_request, mcp_state_request, request};
|
||||
pub(crate) use request::{
|
||||
await_read_request, edit_read_request, json_object_to_prost, mcp_meta_request,
|
||||
};
|
||||
pub use response::{client_event, ClientExecEvent};
|
||||
@@ -0,0 +1,493 @@
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
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" => Message::ShellStreamArgs(pb::ShellArgs {
|
||||
command: string("command")?,
|
||||
working_directory: optional_string("working_directory").unwrap_or_default(),
|
||||
timeout: shell_timeout(call)?,
|
||||
tool_call_id: call.call_id.clone(),
|
||||
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",
|
||||
)?,
|
||||
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(crate) fn await_read_request(
|
||||
id: u32,
|
||||
call: &ToolCall,
|
||||
context: &ExecContext,
|
||||
) -> Result<pb::AgentServerMessage> {
|
||||
let task_id = call
|
||||
.arguments
|
||||
.get("shell_id")
|
||||
.and_then(Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol("AwaitShell is missing shell_id".into()))?;
|
||||
Ok(server_message(
|
||||
id,
|
||||
call,
|
||||
pb::exec_server_message::Message::ReadArgs(pb::ReadArgs {
|
||||
path: format!(
|
||||
"{}/{}.txt",
|
||||
context.terminals_folder.trim_end_matches('/'),
|
||||
task_id
|
||||
),
|
||||
tool_call_id: call.call_id.clone(),
|
||||
..Default::default()
|
||||
}),
|
||||
Some(false),
|
||||
))
|
||||
}
|
||||
|
||||
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_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) }
|
||||
}
|
||||
@@ -0,0 +1,357 @@
|
||||
use crate::{
|
||||
cursor::{
|
||||
interaction,
|
||||
proto::agent::v1 as pb,
|
||||
tools::{
|
||||
edit,
|
||||
result::{self, ToolCompletion},
|
||||
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
||||
},
|
||||
},
|
||||
model::ToolCall,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::request::{await_read_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> {
|
||||
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::Await(_) => advance_await(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)
|
||||
}
|
||||
|
||||
async fn advance_await(
|
||||
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("AwaitShell expected ReadResult".into())),
|
||||
};
|
||||
let ExecStage::Await(state) = &entry.stage else {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell result reached a non-await execution stage".into(),
|
||||
));
|
||||
};
|
||||
let content = match read.result.as_ref() {
|
||||
Some(pb::read_result::Result::Success(success)) => match success.output.as_ref() {
|
||||
Some(pb::read_success::Output::Content(content)) => content.as_str(),
|
||||
_ => "",
|
||||
},
|
||||
Some(pb::read_result::Result::FileNotFound(_)) => "",
|
||||
Some(pb::read_result::Result::Error(error)) => {
|
||||
return Ok(ClientExecEvent::Completed(Box::new(result::await_error(
|
||||
entry,
|
||||
&error.error,
|
||||
)?)))
|
||||
}
|
||||
_ => "",
|
||||
};
|
||||
let regex_match = state
|
||||
.regex
|
||||
.as_ref()
|
||||
.map(|pattern| regex::Regex::new(pattern))
|
||||
.transpose()
|
||||
.map_err(|error| Error::Protocol(format!("invalid AwaitShell pattern: {error}")))?
|
||||
.and_then(|pattern| {
|
||||
pattern
|
||||
.find(content)
|
||||
.map(|found| found.as_str().to_string())
|
||||
});
|
||||
let exit_code = content.lines().find_map(|line| {
|
||||
line.strip_prefix("exit_code:")
|
||||
.and_then(|value| value.trim().parse::<i32>().ok())
|
||||
});
|
||||
if regex_match.is_some() || exit_code.is_some() || std::time::Instant::now() >= state.deadline {
|
||||
return Ok(ClientExecEvent::Completed(Box::new(result::await_result(
|
||||
entry,
|
||||
content.len() as u64,
|
||||
regex_match,
|
||||
exit_code,
|
||||
)?)));
|
||||
}
|
||||
let state = match entry.stage {
|
||||
ExecStage::Await(state) => state,
|
||||
_ => {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell result changed execution stage".into(),
|
||||
))
|
||||
}
|
||||
};
|
||||
let wait = state
|
||||
.deadline
|
||||
.saturating_duration_since(std::time::Instant::now())
|
||||
.min(std::time::Duration::from_secs(1));
|
||||
tokio::time::sleep(wait).await;
|
||||
let call = entry.call.clone();
|
||||
let context = entry.context.clone();
|
||||
let id = registry
|
||||
.reserve_await_again(&call, &context, state, entry.started_at_ms)
|
||||
.await?;
|
||||
Ok(ClientExecEvent::Message(Box::new(await_read_request(
|
||||
id, &call, &context,
|
||||
)?)))
|
||||
}
|
||||
|
||||
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(),
|
||||
})
|
||||
};
|
||||
interaction::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(),
|
||||
},
|
||||
)))
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
//! AwaitShell's timed and file-backed execution paths.
|
||||
|
||||
use crate::{model::ToolCall, Error, Result};
|
||||
|
||||
use super::ToolStart;
|
||||
use crate::cursor::tools::{
|
||||
codec, result,
|
||||
result::ToolResultSender,
|
||||
runtime::{CursorToolRuntime, ExecContext},
|
||||
};
|
||||
|
||||
pub(super) async fn start(
|
||||
runtime: &CursorToolRuntime,
|
||||
results: &ToolResultSender,
|
||||
call: &ToolCall,
|
||||
context: &ExecContext,
|
||||
) -> Result<ToolStart> {
|
||||
let message = if call
|
||||
.arguments
|
||||
.get("shell_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some()
|
||||
{
|
||||
let id = runtime.reserve_await(call, context).await?;
|
||||
Some(codec::await_read_request(id, call, context)?)
|
||||
} else {
|
||||
wait_without_shell_id(results, call)?;
|
||||
None
|
||||
};
|
||||
Ok(ToolStart {
|
||||
messages: message.into_iter().collect(),
|
||||
completion: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn wait_without_shell_id(results: &ToolResultSender, call: &ToolCall) -> Result<()> {
|
||||
let block_ms = call
|
||||
.arguments
|
||||
.get("block_until_ms")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(30_000);
|
||||
if block_ms == 0 || block_ms > 7_140_000 {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell without shell_id requires block_until_ms in 1..=7140000".into(),
|
||||
));
|
||||
}
|
||||
let call = call.clone();
|
||||
let results = results.clone();
|
||||
tokio::spawn(async move {
|
||||
tokio::time::sleep(std::time::Duration::from_millis(block_ms)).await;
|
||||
results.send(result::await_sleep(&call, block_ms));
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
//! 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,71 @@
|
||||
//! Direct Exec and dynamic MCP dispatch.
|
||||
|
||||
use crate::{cursor::proto::agent::v1 as pb, model::ToolCall, Error, Result};
|
||||
|
||||
use super::{normalized, ToolStart};
|
||||
use crate::cursor::tools::{
|
||||
codec, result,
|
||||
runtime::{CursorToolRuntime, ExecContext},
|
||||
};
|
||||
|
||||
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,258 @@
|
||||
//! Interaction query dispatch and approval continuation.
|
||||
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
model::ToolCall,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{normalized, InteractionContinuation, ToolStart};
|
||||
use crate::cursor::tools::{
|
||||
result::{self, ToolResultSender},
|
||||
runtime::{CursorToolRuntime, PendingInteraction},
|
||||
};
|
||||
|
||||
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(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::{response::Html, routing::get, Router};
|
||||
use serde_json::json;
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
use crate::{
|
||||
cursor::{proto::agent::v1 as pb, tools::result::tool_result_channel},
|
||||
model::ToolCall,
|
||||
web::{HtmlEngine, WebFetch, WebSearch},
|
||||
};
|
||||
|
||||
use super::{resume, InteractionContinuation, PendingInteraction};
|
||||
|
||||
#[tokio::test]
|
||||
async fn approved_web_search_completes_through_the_async_result_channel() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move {
|
||||
axum::serve(
|
||||
listener,
|
||||
Router::new().route(
|
||||
"/search",
|
||||
get(|| async {
|
||||
Html(
|
||||
r#"<div class="result"><a class="title" href="https://example.com">Example</a><p class="snippet">Result</p></div>"#,
|
||||
)
|
||||
}),
|
||||
),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
});
|
||||
let search = WebSearch::with_engines(vec![HtmlEngine::new(
|
||||
"fixture",
|
||||
format!("http://{address}/search?q={{query}}"),
|
||||
".result",
|
||||
".title",
|
||||
"a.title",
|
||||
".snippet",
|
||||
)]);
|
||||
let (sender, mut receiver) = tool_result_channel();
|
||||
let continuation = resume(
|
||||
&sender,
|
||||
&search,
|
||||
&WebFetch::for_test(),
|
||||
pending(),
|
||||
&approved(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(continuation, InteractionContinuation::Pending));
|
||||
let completion = receiver.recv().await.unwrap().unwrap();
|
||||
assert!(!completion.result().is_error);
|
||||
assert!(completion.result().content.contains("https://example.com"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn approved_web_fetch_completes_without_client_exec() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move {
|
||||
axum::serve(
|
||||
listener,
|
||||
Router::new().route(
|
||||
"/article",
|
||||
get(|| async {
|
||||
Html(
|
||||
r#"<html><head><title>Fetched page</title></head><body><article><h1>Fetched page</h1><p>This readable article is long enough for deterministic extraction by the server-side fetch tool.</p><p>It completes directly through the ToolResult channel without creating a Cursor FetchArgs message.</p></article></body></html>"#,
|
||||
)
|
||||
}),
|
||||
),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
});
|
||||
let (sender, mut receiver) = tool_result_channel();
|
||||
let continuation = resume(
|
||||
&sender,
|
||||
&WebSearch::built_in(),
|
||||
&WebFetch::for_test(),
|
||||
pending_fetch(format!("http://{address}/article")),
|
||||
&approved_fetch(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(continuation, InteractionContinuation::Pending));
|
||||
let completion = receiver.recv().await.unwrap().unwrap();
|
||||
assert!(!completion.result().is_error);
|
||||
assert!(completion.result().content.contains("Fetched page"));
|
||||
}
|
||||
|
||||
fn pending() -> PendingInteraction {
|
||||
PendingInteraction {
|
||||
call: ToolCall {
|
||||
index: 0,
|
||||
call_id: "search".into(),
|
||||
model_call_id: "model".into(),
|
||||
name: "WebSearch".into(),
|
||||
arguments_text: r#"{"search_term":"rust"}"#.into(),
|
||||
arguments: json!({"search_term": "rust"}),
|
||||
},
|
||||
started_at_ms: 1,
|
||||
}
|
||||
}
|
||||
|
||||
fn approved() -> pb::InteractionResponse {
|
||||
pb::InteractionResponse {
|
||||
id: 1,
|
||||
result: Some(pb::interaction_response::Result::WebSearchRequestResponse(
|
||||
pb::WebSearchRequestResponse {
|
||||
result: Some(pb::web_search_request_response::Result::Approved(
|
||||
pb::web_search_request_response::Approved::default(),
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn pending_fetch(url: String) -> PendingInteraction {
|
||||
PendingInteraction {
|
||||
call: ToolCall {
|
||||
index: 0,
|
||||
call_id: "fetch".into(),
|
||||
model_call_id: "model".into(),
|
||||
name: "WebFetch".into(),
|
||||
arguments_text: serde_json::to_string(&json!({"url": url})).unwrap(),
|
||||
arguments: json!({"url": url}),
|
||||
},
|
||||
started_at_ms: 1,
|
||||
}
|
||||
}
|
||||
|
||||
fn approved_fetch() -> pb::InteractionResponse {
|
||||
pb::InteractionResponse {
|
||||
id: 2,
|
||||
result: Some(pb::interaction_response::Result::WebFetchRequestResponse(
|
||||
pb::WebFetchRequestResponse {
|
||||
result: Some(pb::web_fetch_request_response::Result::Approved(
|
||||
pb::web_fetch_request_response::Approved::default(),
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
//! Synchronous local tool dispatch.
|
||||
|
||||
use crate::{model::ToolCall, Result};
|
||||
|
||||
use super::ToolStart;
|
||||
use crate::cursor::tools::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,89 @@
|
||||
mod await_shell;
|
||||
mod edit;
|
||||
mod exec;
|
||||
mod interaction;
|
||||
mod local;
|
||||
mod semble;
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::ToolCall,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{
|
||||
result::{ToolCompletion, ToolResultSender},
|
||||
runtime::{CursorToolRuntime, ExecContext, PendingInteraction},
|
||||
};
|
||||
|
||||
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,
|
||||
) -> 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);
|
||||
}
|
||||
|
||||
match normalized(&call.name).as_str() {
|
||||
"shell" | "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),
|
||||
"awaitshell" => await_shell::start(runtime, results, call, context).await,
|
||||
"semblesearch" | "semblefindrelated" => semble::start(results, call),
|
||||
_ => Err(Error::Protocol(format!("unsupported tool: {}", call.name))),
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
pub(super) fn normalized(name: &str) -> String {
|
||||
name.chars()
|
||||
.filter(|character| character.is_ascii_alphanumeric())
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
//! Asynchronous dispatch for the application-owned Semble search tools.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use semble_core::{ContentType, FindRelatedRequest, SearchEngine, SearchRequest, SembleConfig};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
use crate::{model::ToolCall, Error, Result};
|
||||
|
||||
use super::ToolStart;
|
||||
use crate::cursor::tools::{
|
||||
result::{self, ToolResultSender},
|
||||
runtime::now_ms,
|
||||
};
|
||||
|
||||
static ENGINE: OnceCell<Arc<SearchEngine>> = OnceCell::const_new();
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
enum ContentSelection {
|
||||
#[default]
|
||||
Code,
|
||||
Docs,
|
||||
Config,
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct SearchArguments {
|
||||
query: String,
|
||||
repo: String,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct FindRelatedArguments {
|
||||
repo: String,
|
||||
file_path: String,
|
||||
line: usize,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
pub(super) fn start(results: &ToolResultSender, call: &ToolCall) -> Result<ToolStart> {
|
||||
let operation = match super::normalized(&call.name).as_str() {
|
||||
"semblesearch" => Operation::Search(serde_json::from_value(call.arguments.clone())?),
|
||||
"semblefindrelated" => {
|
||||
Operation::FindRelated(serde_json::from_value(call.arguments.clone())?)
|
||||
}
|
||||
_ => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unsupported Semble tool: {}",
|
||||
call.name
|
||||
)))
|
||||
}
|
||||
};
|
||||
let call = call.clone();
|
||||
let results = results.clone();
|
||||
let started_at_ms = now_ms();
|
||||
tokio::spawn(async move {
|
||||
let output = execute(operation).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,
|
||||
})
|
||||
}
|
||||
|
||||
enum Operation {
|
||||
Search(SearchArguments),
|
||||
FindRelated(FindRelatedArguments),
|
||||
}
|
||||
|
||||
async fn execute(operation: Operation) -> std::result::Result<Value, String> {
|
||||
let engine = engine().await.map_err(|error| error.to_string())?;
|
||||
tokio::task::spawn_blocking(move || match operation {
|
||||
Operation::Search(arguments) => engine
|
||||
.search(SearchRequest {
|
||||
query: arguments.query,
|
||||
repo: arguments.repo.into(),
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
Operation::FindRelated(arguments) => engine
|
||||
.find_related(FindRelatedRequest {
|
||||
repo: arguments.repo.into(),
|
||||
file_path: arguments.file_path,
|
||||
line: arguments.line,
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
})
|
||||
.await
|
||||
.map_err(|error| format!("Semble search worker failed: {error}"))?
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
async fn engine() -> Result<Arc<SearchEngine>> {
|
||||
ENGINE
|
||||
.get_or_try_init(|| async {
|
||||
tokio::task::spawn_blocking(|| SearchEngine::load_default(SembleConfig::default()))
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))?
|
||||
.map(Arc::new)
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))
|
||||
})
|
||||
.await
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn json_value(response: semble_core::SearchResponse) -> semble_core::Result<serde_json::Value> {
|
||||
serde_json::to_value(response)
|
||||
.map_err(|error| semble_core::Error::Serialization(error.to_string()))
|
||||
}
|
||||
|
||||
fn content(selection: ContentSelection) -> Vec<ContentType> {
|
||||
match selection {
|
||||
ContentSelection::Code => vec![ContentType::Code],
|
||||
ContentSelection::Docs => vec![ContentType::Docs],
|
||||
ContentSelection::Config => vec![ContentType::Config],
|
||||
ContentSelection::All => vec![ContentType::Code, ContentType::Docs, ContentType::Config],
|
||||
}
|
||||
}
|
||||
|
||||
fn default_top_k() -> usize {
|
||||
5
|
||||
}
|
||||
|
||||
fn default_snippet_lines() -> Option<usize> {
|
||||
Some(10)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn search_arguments_use_code_search_defaults() {
|
||||
let arguments: SearchArguments = serde_json::from_value(json!({
|
||||
"query": "request persistence",
|
||||
"repo": "/tmp/repo"
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.top_k, 5);
|
||||
assert_eq!(arguments.max_snippet_lines, Some(10));
|
||||
assert!(matches!(arguments.content, ContentSelection::Code));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn find_related_does_not_require_a_ui_description() {
|
||||
let arguments: FindRelatedArguments = serde_json::from_value(json!({
|
||||
"repo": "/tmp/repo",
|
||||
"file_path": "src/auth.ts",
|
||||
"line": 42
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.file_path, "src/auth.ts");
|
||||
assert_eq!(arguments.line, 42);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn all_content_expands_to_every_indexed_scope() {
|
||||
assert_eq!(
|
||||
content(ContentSelection::All),
|
||||
vec![ContentType::Code, ContentType::Docs, ContentType::Config]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
use serde_json::Value;
|
||||
use similar::{ChangeTag, TextDiff};
|
||||
|
||||
use crate::{model::ToolCall, Error, Result};
|
||||
|
||||
use crate::cursor::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 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(¬ebook)
|
||||
.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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn call(name: &str, arguments: Value) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "call\nfc_1".into(),
|
||||
model_call_id: "model".into(),
|
||||
name: name.into(),
|
||||
arguments_text: String::new(),
|
||||
arguments,
|
||||
}
|
||||
}
|
||||
|
||||
fn read(content: &str) -> pb::ReadResult {
|
||||
pb::ReadResult {
|
||||
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
|
||||
output: Some(pb::read_success::Output::Content(content.into())),
|
||||
..Default::default()
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn write_and_str_replace_use_one_lf_canonical_form() {
|
||||
let write = after_read(
|
||||
&call("Write", json!({"path":"/a","contents":"new\rline\r\n"})),
|
||||
&read("old\r\nline\r"),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(write.before, "old\nline\n");
|
||||
assert_eq!(write.after, "new\nline\n");
|
||||
|
||||
let replacement = after_read(
|
||||
&call(
|
||||
"StrReplace",
|
||||
json!({"path":"/a","old_string":"old\nline","new_string":"new\r\nline"}),
|
||||
),
|
||||
&read("old\r\nline\r\nrest"),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(replacement.after, "new\nline\nrest");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn str_replace_requires_one_match_unless_replace_all_is_explicit() {
|
||||
let ambiguous = after_read(
|
||||
&call(
|
||||
"StrReplace",
|
||||
json!({"path":"/a","old_string":"same","new_string":"new"}),
|
||||
),
|
||||
&read("same\nsame\n"),
|
||||
)
|
||||
.unwrap_err();
|
||||
assert_eq!(ambiguous, "old_string is not unique; found 2 occurrences");
|
||||
|
||||
let all = after_read(
|
||||
&call(
|
||||
"StrReplace",
|
||||
json!({
|
||||
"path":"/a", "old_string":"same", "new_string":"new",
|
||||
"replace_all":true
|
||||
}),
|
||||
),
|
||||
&read("same\rsame\r\n"),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(all.after, "new\nnew\n");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn notebook_edit_targets_one_cell_and_preserves_lf() {
|
||||
let notebook = r#"{"cells":[{"cell_type":"code","source":["old\r\n","line"]}],"metadata":{},"nbformat":4,"nbformat_minor":5}"#;
|
||||
let edit = after_read(
|
||||
&call(
|
||||
"EditNotebook",
|
||||
json!({
|
||||
"target_notebook":"/a.ipynb", "cell_idx":0, "is_new_cell":false,
|
||||
"cell_language":"python", "old_string":"old\nline", "new_string":"new\r\nline"
|
||||
}),
|
||||
),
|
||||
&read(notebook),
|
||||
)
|
||||
.unwrap();
|
||||
let parsed: Value = serde_json::from_str(&edit.after).unwrap();
|
||||
assert_eq!(parsed["cells"][0]["source"], json!(["new\n", "line"]));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
use std::collections::{BTreeMap, HashSet};
|
||||
|
||||
pub mod codec;
|
||||
mod dispatch;
|
||||
pub(crate) mod edit;
|
||||
pub(crate) mod result;
|
||||
pub mod runtime;
|
||||
pub(crate) mod stream;
|
||||
|
||||
use crate::{
|
||||
model::{CanonicalMessage, MessageContent, Role, ToolCall},
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use self::result::{ToolCompletion, ToolResultSender};
|
||||
use super::{interaction, proto::agent::v1 as pb};
|
||||
use runtime::{CursorToolRuntime, ExecContext};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ToolDispatcher {
|
||||
runtime: CursorToolRuntime,
|
||||
results: ToolResultSender,
|
||||
search: WebSearch,
|
||||
fetch: WebFetch,
|
||||
}
|
||||
|
||||
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, _) = result::tool_result_channel();
|
||||
Self::with_results(runtime, results)
|
||||
}
|
||||
|
||||
pub fn with_results(runtime: CursorToolRuntime, results: ToolResultSender) -> Self {
|
||||
Self {
|
||||
runtime,
|
||||
results,
|
||||
search: WebSearch::built_in(),
|
||||
fetch: WebFetch::built_in(),
|
||||
}
|
||||
}
|
||||
|
||||
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;
|
||||
}
|
||||
dispatched.push(
|
||||
self.start(
|
||||
call,
|
||||
first_tool_index + position,
|
||||
!state.started.contains(&call.call_id),
|
||||
dynamic_mcp,
|
||||
context,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
Ok(dispatched)
|
||||
}
|
||||
|
||||
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![interaction::tool_started(
|
||||
&call,
|
||||
dynamic_mcp.get(&call.name),
|
||||
)?]
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
let started = dispatch::start(
|
||||
&self.runtime,
|
||||
&self.results,
|
||||
&call,
|
||||
message_index,
|
||||
dynamic_mcp,
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
messages.extend(started.messages);
|
||||
Ok(DispatchedTool {
|
||||
messages,
|
||||
completion: started.completion,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn interaction_response(
|
||||
&self,
|
||||
response: &pb::InteractionResponse,
|
||||
) -> Result<ClientToolEvent> {
|
||||
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 dispatch::resume_interaction(
|
||||
&self.results,
|
||||
&self.search,
|
||||
&self.fetch,
|
||||
pending,
|
||||
response,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
dispatch::InteractionContinuation::Completed(completion) => {
|
||||
ClientToolEvent::Completed(completion)
|
||||
}
|
||||
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()
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::{ToolCall, ToolResult},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{now_ms, ToolCompletion};
|
||||
use crate::cursor::tools::runtime::{ExecStage, PendingExec};
|
||||
|
||||
pub(crate) fn await_result(
|
||||
pending: PendingExec,
|
||||
output_length: u64,
|
||||
regex_match: Option<String>,
|
||||
exit_code: Option<i32>,
|
||||
) -> Result<ToolCompletion> {
|
||||
let ExecStage::Await(state) = &pending.stage else {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell completion reached a non-await execution stage".into(),
|
||||
));
|
||||
};
|
||||
let runtime_ms = now_ms().saturating_sub(pending.started_at_ms);
|
||||
let result = if exit_code.is_some() {
|
||||
pb::await_success::AwaitResult::Complete(pb::AwaitTaskComplete {
|
||||
task_id: state.task_id.clone(),
|
||||
runtime_ms,
|
||||
output_file_path: state.output_file_path.clone(),
|
||||
output_length,
|
||||
regex_requested: state.regex.is_some(),
|
||||
regex_match,
|
||||
exit_code,
|
||||
wake_reason: Some("task_complete".into()),
|
||||
})
|
||||
} else {
|
||||
pb::await_success::AwaitResult::StillRunning(pb::AwaitTaskStillRunning {
|
||||
task_id: state.task_id.clone(),
|
||||
runtime_ms,
|
||||
output_file_path: state.output_file_path.clone(),
|
||||
output_length,
|
||||
regex_requested: state.regex.is_some(),
|
||||
regex_match,
|
||||
wake_reason: Some("timeout_or_pattern".into()),
|
||||
})
|
||||
};
|
||||
let content = serde_json::json!({
|
||||
"task_id": state.task_id,
|
||||
"output_file_path": state.output_file_path,
|
||||
"output_length": output_length,
|
||||
"exit_code": exit_code,
|
||||
})
|
||||
.to_string();
|
||||
completion(
|
||||
&pending,
|
||||
content,
|
||||
false,
|
||||
pb::await_result::Result::Success(pb::AwaitSuccess {
|
||||
await_result: Some(result),
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn await_error(pending: PendingExec, error: &str) -> Result<ToolCompletion> {
|
||||
completion(
|
||||
&pending,
|
||||
error.into(),
|
||||
true,
|
||||
pb::await_result::Result::Error(pb::AwaitError {
|
||||
error: error.into(),
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn completion(
|
||||
pending: &PendingExec,
|
||||
content: String,
|
||||
is_error: bool,
|
||||
result: pb::await_result::Result,
|
||||
) -> Result<ToolCompletion> {
|
||||
let ExecStage::Await(state) = &pending.stage else {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell completion reached a non-await execution stage".into(),
|
||||
));
|
||||
};
|
||||
Ok(ToolCompletion::new(
|
||||
&pending.call,
|
||||
pending.started_at_ms,
|
||||
ToolResult {
|
||||
call_id: pending.call.call_id.clone(),
|
||||
content,
|
||||
is_error,
|
||||
image: None,
|
||||
},
|
||||
pb::tool_call::Tool::AwaitToolCall(pb::AwaitToolCall {
|
||||
args: Some(pb::AwaitArgs {
|
||||
task_id: state.task_id.clone(),
|
||||
block_until_ms: pending
|
||||
.call
|
||||
.arguments
|
||||
.get("block_until_ms")
|
||||
.and_then(Value::as_u64)
|
||||
.map(|value| value as u32),
|
||||
regex: state.regex.clone(),
|
||||
}),
|
||||
result: Some(pb::AwaitResult {
|
||||
result: Some(result),
|
||||
}),
|
||||
}),
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn await_sleep(call: &ToolCall, runtime_ms: u64) -> ToolCompletion {
|
||||
ToolCompletion::new(
|
||||
call,
|
||||
now_ms().saturating_sub(runtime_ms),
|
||||
ToolResult {
|
||||
call_id: call.call_id.clone(),
|
||||
content: format!("Waited {runtime_ms} ms"),
|
||||
is_error: false,
|
||||
image: None,
|
||||
},
|
||||
pb::tool_call::Tool::AwaitToolCall(pb::AwaitToolCall {
|
||||
args: Some(pb::AwaitArgs {
|
||||
task_id: String::new(),
|
||||
block_until_ms: Some(runtime_ms as u32),
|
||||
regex: None,
|
||||
}),
|
||||
result: Some(pb::AwaitResult {
|
||||
result: Some(pb::await_result::Result::Success(pb::AwaitSuccess {
|
||||
await_result: Some(pb::await_success::AwaitResult::StillRunning(
|
||||
pb::AwaitTaskStillRunning {
|
||||
task_id: String::new(),
|
||||
runtime_ms,
|
||||
output_file_path: String::new(),
|
||||
output_length: 0,
|
||||
regex_requested: false,
|
||||
regex_match: None,
|
||||
wake_reason: Some("sleep_complete".into()),
|
||||
},
|
||||
)),
|
||||
})),
|
||||
}),
|
||||
}),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
mod output;
|
||||
mod render;
|
||||
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
model::ToolResult,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{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};
|
||||
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),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use base64::{engine::general_purpose::STANDARD, Engine};
|
||||
|
||||
use super::image_mime_type;
|
||||
|
||||
#[test]
|
||||
fn read_image_requires_a_decodable_supported_image() {
|
||||
let png = STANDARD
|
||||
.decode("iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=")
|
||||
.unwrap();
|
||||
assert_eq!(image_mime_type(&png), Some("image/png"));
|
||||
assert_eq!(image_mime_type(b"\x89PNG\r\n\x1a\n"), None);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
use crate::{cursor::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((
|
||||
format!(
|
||||
"grep success pattern={} mode={}",
|
||||
value.pattern, value.output_mode
|
||||
),
|
||||
false,
|
||||
)),
|
||||
R::Error(value) => Ok((value.error.clone(), true)),
|
||||
}
|
||||
}
|
||||
|
||||
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((
|
||||
format!(
|
||||
"diagnostics path={} count={}",
|
||||
value.path, value.total_diagnostics
|
||||
),
|
||||
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 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,254 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{cursor::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,471 @@
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
web::{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"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::ToolCall,
|
||||
web::{FetchedPage, SearchHit},
|
||||
};
|
||||
|
||||
use super::{complete_web_fetch, complete_web_search, PendingInteraction};
|
||||
|
||||
#[test]
|
||||
fn web_search_success_becomes_a_typed_tool_result() {
|
||||
let completion = complete_web_search(
|
||||
pending(),
|
||||
Ok(vec![SearchHit::new(
|
||||
"Rust",
|
||||
"https://www.rust-lang.org",
|
||||
"A language empowering everyone",
|
||||
vec!["first", "second"],
|
||||
)]),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(!completion.result().is_error);
|
||||
assert!(completion
|
||||
.result()
|
||||
.content
|
||||
.contains("https://www.rust-lang.org"));
|
||||
let Some(pb::tool_call::Tool::WebSearchToolCall(tool)) =
|
||||
completion.tool_call().tool.as_ref()
|
||||
else {
|
||||
panic!("expected WebSearchToolCall")
|
||||
};
|
||||
let Some(pb::web_search_result::Result::Success(success)) = tool
|
||||
.result
|
||||
.as_ref()
|
||||
.and_then(|result| result.result.as_ref())
|
||||
else {
|
||||
panic!("expected WebSearchSuccess")
|
||||
};
|
||||
assert_eq!(success.references.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn web_search_failure_is_a_tool_error_instead_of_a_run_error() {
|
||||
let completion = complete_web_search(pending(), Err("all engines failed".into())).unwrap();
|
||||
|
||||
assert!(completion.result().is_error);
|
||||
assert_eq!(completion.result().content, "all engines failed");
|
||||
let Some(pb::tool_call::Tool::WebSearchToolCall(tool)) =
|
||||
completion.tool_call().tool.as_ref()
|
||||
else {
|
||||
panic!("expected WebSearchToolCall")
|
||||
};
|
||||
assert!(matches!(
|
||||
tool.result
|
||||
.as_ref()
|
||||
.and_then(|result| result.result.as_ref()),
|
||||
Some(pb::web_search_result::Result::Error(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn web_fetch_success_becomes_markdown_tool_result() {
|
||||
let completion = complete_web_fetch(
|
||||
pending_fetch(),
|
||||
Ok(FetchedPage {
|
||||
url: "https://example.com/final".into(),
|
||||
markdown: "# Article\n\nReadable body.".into(),
|
||||
}),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(!completion.result().is_error);
|
||||
assert_eq!(completion.result().content, "# Article\n\nReadable body.");
|
||||
let Some(pb::tool_call::Tool::WebFetchToolCall(tool)) =
|
||||
completion.tool_call().tool.as_ref()
|
||||
else {
|
||||
panic!("expected WebFetchToolCall")
|
||||
};
|
||||
let Some(pb::web_fetch_result::Result::Success(success)) = tool
|
||||
.result
|
||||
.as_ref()
|
||||
.and_then(|result| result.result.as_ref())
|
||||
else {
|
||||
panic!("expected WebFetchSuccess")
|
||||
};
|
||||
assert_eq!(success.url, "https://example.com/final");
|
||||
assert_eq!(success.markdown, "# Article\n\nReadable body.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn web_fetch_failure_is_a_tool_error_instead_of_a_run_error() {
|
||||
let completion =
|
||||
complete_web_fetch(pending_fetch(), Err("unsupported content type".into())).unwrap();
|
||||
|
||||
assert!(completion.result().is_error);
|
||||
assert_eq!(completion.result().content, "unsupported content type");
|
||||
let Some(pb::tool_call::Tool::WebFetchToolCall(tool)) =
|
||||
completion.tool_call().tool.as_ref()
|
||||
else {
|
||||
panic!("expected WebFetchToolCall")
|
||||
};
|
||||
assert!(matches!(
|
||||
tool.result
|
||||
.as_ref()
|
||||
.and_then(|result| result.result.as_ref()),
|
||||
Some(pb::web_fetch_result::Result::Error(_))
|
||||
));
|
||||
}
|
||||
|
||||
fn pending() -> PendingInteraction {
|
||||
PendingInteraction {
|
||||
call: ToolCall {
|
||||
index: 0,
|
||||
call_id: "search-call".into(),
|
||||
model_call_id: "model-call".into(),
|
||||
name: "WebSearch".into(),
|
||||
arguments_text: r#"{"search_term":"rust"}"#.into(),
|
||||
arguments: json!({"search_term": "rust"}),
|
||||
},
|
||||
started_at_ms: 1,
|
||||
}
|
||||
}
|
||||
|
||||
fn pending_fetch() -> PendingInteraction {
|
||||
PendingInteraction {
|
||||
call: ToolCall {
|
||||
index: 0,
|
||||
call_id: "fetch-call".into(),
|
||||
model_call_id: "model-call".into(),
|
||||
name: "WebFetch".into(),
|
||||
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
||||
arguments: json!({"url": "https://example.com"}),
|
||||
},
|
||||
started_at_ms: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn disabled_subagent_returns_a_model_visible_system_reminder() {
|
||||
let arguments = serde_json::json!({
|
||||
"description":"inspect",
|
||||
"prompt":"inspect",
|
||||
"subagent_type":"explore"
|
||||
});
|
||||
let completion = subagents_disabled(&ToolCall {
|
||||
index: 0,
|
||||
call_id: "task-1".into(),
|
||||
model_call_id: "model-call-1".into(),
|
||||
name: "Task".into(),
|
||||
arguments_text: arguments.to_string(),
|
||||
arguments,
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(completion.result().content, SUBAGENTS_DISABLED_REMINDER);
|
||||
assert!(completion.result().is_error);
|
||||
assert!(matches!(
|
||||
completion.tool_call().tool.as_ref(),
|
||||
Some(pb::tool_call::Tool::TaskToolCall(pb::TaskToolCall {
|
||||
result: Some(pb::TaskResult {
|
||||
result: Some(pb::task_result::Result::Error(pb::TaskError { error })),
|
||||
}),
|
||||
..
|
||||
})) if error == SUBAGENTS_DISABLED_REMINDER
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
fn normalized(name: &str) -> String {
|
||||
name.chars()
|
||||
.filter(|character| character.is_ascii_alphanumeric())
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
//! Canonical failures produced before an MCP request reaches the Cursor client.
|
||||
|
||||
use crate::{
|
||||
cursor::{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,143 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{cursor::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,173 @@
|
||||
mod await_shell;
|
||||
mod exec;
|
||||
mod interaction;
|
||||
mod local;
|
||||
mod mcp;
|
||||
mod mcp_state;
|
||||
mod semble;
|
||||
|
||||
use serde_json::Value;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::{ToolCall, ToolImageReference, ToolResult},
|
||||
store::BlobId,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::runtime::now_ms;
|
||||
|
||||
pub(crate) use await_shell::{await_error, await_result, await_sleep};
|
||||
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 semble::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,
|
||||
result: ToolResult,
|
||||
tool: pb::tool_call::Tool,
|
||||
) -> Self {
|
||||
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,144 @@
|
||||
//! Cursor MCP-card rendering for direct Semble Agent tools.
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
cursor::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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn direct_search_renders_as_a_builtin_semble_mcp_card() {
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: "SembleSearch".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"description": "Find request tracing",
|
||||
"repo": "/tmp/repo",
|
||||
"query": "request tracing"
|
||||
}),
|
||||
};
|
||||
let completion = complete(&call, 1, Ok(json!({"results": []}))).unwrap();
|
||||
let pb::tool_call::Tool::McpToolCall(tool) = completion.tool_call().tool.as_ref().unwrap()
|
||||
else {
|
||||
panic!("expected MCP tool card");
|
||||
};
|
||||
assert_eq!(tool.description.as_deref(), Some("Find request tracing"));
|
||||
let args = tool.args.as_ref().unwrap();
|
||||
assert_eq!(args.provider_identifier, PROVIDER_IDENTIFIER);
|
||||
assert_eq!(args.tool_name, "search");
|
||||
assert!(!args.args.contains_key("description"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,429 @@
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::{
|
||||
atomic::{AtomicU32, Ordering},
|
||||
Arc,
|
||||
},
|
||||
time::Instant,
|
||||
};
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::{cursor::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>>>,
|
||||
}
|
||||
|
||||
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),
|
||||
Await(AwaitState),
|
||||
}
|
||||
|
||||
pub(crate) struct AwaitState {
|
||||
pub deadline: Instant,
|
||||
pub output_file_path: String,
|
||||
pub task_id: String,
|
||||
pub regex: Option<String>,
|
||||
}
|
||||
|
||||
#[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 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
|
||||
}
|
||||
|
||||
pub(crate) async fn reserve_await(
|
||||
&self,
|
||||
call: &ToolCall,
|
||||
context: &ExecContext,
|
||||
) -> Result<u32> {
|
||||
let task_id = call
|
||||
.arguments
|
||||
.get("shell_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol("AwaitShell is missing shell_id".into()))?;
|
||||
let block_ms = call
|
||||
.arguments
|
||||
.get("block_until_ms")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(30_000);
|
||||
if block_ms > 7_140_000 {
|
||||
return Err(Error::Protocol(
|
||||
"AwaitShell block_until_ms exceeds 7140000".into(),
|
||||
));
|
||||
}
|
||||
let output_file_path = format!(
|
||||
"{}/{}.txt",
|
||||
context.terminals_folder.trim_end_matches('/'),
|
||||
task_id
|
||||
);
|
||||
self.reserve_exec_stage(
|
||||
call,
|
||||
context,
|
||||
ExecStage::Await(AwaitState {
|
||||
deadline: Instant::now() + std::time::Duration::from_millis(block_ms),
|
||||
output_file_path,
|
||||
task_id: task_id.to_string(),
|
||||
regex: call
|
||||
.arguments
|
||||
.get("pattern")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_string),
|
||||
}),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn reserve_await_again(
|
||||
&self,
|
||||
call: &ToolCall,
|
||||
context: &ExecContext,
|
||||
state: AwaitState,
|
||||
started_at_ms: u64,
|
||||
) -> Result<u32> {
|
||||
self.reserve_exec_stage(call, context, ExecStage::Await(state), 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 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();
|
||||
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
|
||||
}
|
||||
|
||||
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()))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn task(arguments: serde_json::Value) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "task-1".into(),
|
||||
model_call_id: "model-call-1".into(),
|
||||
name: "Task".into(),
|
||||
arguments_text: arguments.to_string(),
|
||||
arguments,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_model_defaults_to_parent_and_honors_an_explicit_model() {
|
||||
let context = ExecContext {
|
||||
default_subagent_model: "parent-model".into(),
|
||||
..ExecContext::default()
|
||||
};
|
||||
let inherited = context
|
||||
.prepare_call(&task(serde_json::json!({"prompt":"inspect"})))
|
||||
.unwrap();
|
||||
let explicit = context
|
||||
.prepare_call(&task(serde_json::json!({
|
||||
"prompt":"inspect",
|
||||
"model":"child-model"
|
||||
})))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(inherited.arguments["model"], "parent-model");
|
||||
assert_eq!(explicit.arguments["model"], "child-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn global_subagent_model_applies_to_every_task_type() {
|
||||
let context = ExecContext {
|
||||
default_subagent_model: "parent-model".into(),
|
||||
subagent_model: Some(SubagentModel::Model("child-model".into())),
|
||||
..ExecContext::default()
|
||||
};
|
||||
let call = task(serde_json::json!({
|
||||
"prompt":"inspect",
|
||||
"subagent_type":"test-subagent"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
context.prepare_call(&call).unwrap().arguments["model"],
|
||||
"child-model"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_subagents_disable_every_task_type() {
|
||||
let context = ExecContext {
|
||||
default_subagent_model: "parent-model".into(),
|
||||
subagent_model: Some(SubagentModel::Disabled),
|
||||
..ExecContext::default()
|
||||
};
|
||||
let call = task(serde_json::json!({
|
||||
"prompt":"inspect",
|
||||
"subagent_type":"test-subagent"
|
||||
}));
|
||||
|
||||
assert!(context.task_disabled(&call));
|
||||
assert!(context
|
||||
.prepare_call(&call)
|
||||
.unwrap()
|
||||
.arguments
|
||||
.get("model")
|
||||
.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn now_ms() -> u64 {
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_millis() as u64
|
||||
}
|
||||
@@ -0,0 +1,341 @@
|
||||
use crate::{
|
||||
cursor::{
|
||||
interaction,
|
||||
json_stream::{JsonStringFields, StringFieldEvent},
|
||||
proto::agent::v1 as pb,
|
||||
},
|
||||
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()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn call(name: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: name.into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: Value::Null,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn plain_tools_only_project_raw_argument_deltas() {
|
||||
let call = call("Read");
|
||||
let mut stream = ToolCallStream::new(&call.name, None);
|
||||
assert_eq!(
|
||||
stream.arguments_delta(&call, "{\"path\":").unwrap().len(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn write_projects_path_and_content_without_starting_execution() {
|
||||
let call = call("Write");
|
||||
let mut stream = ToolCallStream::new(&call.name, None);
|
||||
let first = stream
|
||||
.arguments_delta(&call, "{\"path\":\"/tmp/a\",\"contents\":\"hel")
|
||||
.unwrap();
|
||||
assert_eq!(first.len(), 2);
|
||||
assert!(matches!(
|
||||
first[0].message,
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(
|
||||
pb::InteractionUpdate {
|
||||
message: Some(pb::interaction_update::Message::PartialToolCall(_))
|
||||
}
|
||||
))
|
||||
));
|
||||
assert_eq!(edit_delta(&first[1]), "hel");
|
||||
|
||||
let second = stream.arguments_delta(&call, "lo\\n世界\"}").unwrap();
|
||||
assert_eq!(second.len(), 1);
|
||||
assert_eq!(edit_delta(&second[0]), "lo\n世界");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn str_replace_projects_only_new_string_when_path_arrives_later() {
|
||||
let mut call = call("StrReplace");
|
||||
let mut stream = ToolCallStream::new(&call.name, None);
|
||||
let first = stream
|
||||
.arguments_delta(&call, "{\"new_string\":\"new\",\"old_string\":\"old\",")
|
||||
.unwrap();
|
||||
assert_eq!(first.len(), 1);
|
||||
assert_eq!(edit_delta(&first[0]), "new");
|
||||
let second = stream
|
||||
.arguments_delta(&call, "\"path\":\"/tmp/a\"}")
|
||||
.unwrap();
|
||||
assert_eq!(second.len(), 1);
|
||||
assert!(matches!(
|
||||
second[0].message,
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(
|
||||
pb::InteractionUpdate {
|
||||
message: Some(pb::interaction_update::Message::PartialToolCall(_))
|
||||
}
|
||||
))
|
||||
));
|
||||
|
||||
call.arguments = json!({
|
||||
"path": "/tmp/a",
|
||||
"old_string": "old",
|
||||
"new_string": "new"
|
||||
});
|
||||
let rendered = interaction::render_tool_call(&call, false).unwrap();
|
||||
let Some(pb::tool_call::Tool::EditToolCall(edit)) = rendered.tool else {
|
||||
panic!("expected EditToolCall")
|
||||
};
|
||||
assert_eq!(edit.args.unwrap().stream_content.as_deref(), Some("new"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn edit_stream_normalizes_split_crlf_once() {
|
||||
let call = call("Write");
|
||||
let mut stream = ToolCallStream::new(&call.name, None);
|
||||
let first = stream
|
||||
.arguments_delta(&call, "{\"contents\":\"a\\r")
|
||||
.unwrap();
|
||||
let second = stream
|
||||
.arguments_delta(&call, "\\nb\\r\",\"path\":\"/tmp/a\"}")
|
||||
.unwrap();
|
||||
assert_eq!(edit_delta(&first[0]), "a");
|
||||
assert_eq!(edit_delta(&second[0]), "\nb");
|
||||
assert_eq!(edit_delta(&second[1]), "\n");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_plan_projects_completed_fields_as_structured_partial_args() {
|
||||
let call = call("CreatePlan");
|
||||
let mut stream = ToolCallStream::new(&call.name, None);
|
||||
let name = stream
|
||||
.arguments_delta(&call, "{\"name\":\"Migration Plan\",\"plan\":\"# Move")
|
||||
.unwrap();
|
||||
assert_eq!(name.len(), 1);
|
||||
|
||||
let messages = stream
|
||||
.arguments_delta(
|
||||
&call,
|
||||
" services\",\"overview\":\"Move the services safely\",\"todos\":[]}",
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(messages.len(), 1);
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) =
|
||||
&messages[0].message
|
||||
else {
|
||||
panic!("expected InteractionUpdate")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::PartialToolCall(partial)) = &update.message
|
||||
else {
|
||||
panic!("expected PartialToolCall")
|
||||
};
|
||||
assert!(partial.args_text_delta.is_empty());
|
||||
let Some(pb::tool_call::Tool::CreatePlanToolCall(plan)) = partial
|
||||
.tool_call
|
||||
.as_ref()
|
||||
.and_then(|call| call.tool.as_ref())
|
||||
else {
|
||||
panic!("expected CreatePlanToolCall")
|
||||
};
|
||||
let args = plan.args.as_ref().unwrap();
|
||||
assert_eq!(args.name, "Migration Plan");
|
||||
assert_eq!(args.plan, "# Move services");
|
||||
assert_eq!(args.overview, "Move the services safely");
|
||||
assert!(args.todos.is_empty());
|
||||
}
|
||||
|
||||
fn edit_delta(message: &pb::AgentServerMessage) -> &str {
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = &message.message
|
||||
else {
|
||||
panic!("expected InteractionUpdate")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::ToolCallDelta(update)) = &update.message else {
|
||||
panic!("expected ToolCallDelta")
|
||||
};
|
||||
let Some(pb::tool_call_delta::Delta::EditToolCallDelta(delta)) = update
|
||||
.tool_call_delta
|
||||
.as_deref()
|
||||
.and_then(|delta| delta.delta.as_ref())
|
||||
else {
|
||||
panic!("expected EditToolCallDelta")
|
||||
};
|
||||
&delta.stream_content_delta
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,331 @@
|
||||
use std::collections::HashSet;
|
||||
|
||||
use crate::{
|
||||
cursor::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;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::{CanonicalMessage, Origin, Role};
|
||||
|
||||
#[test]
|
||||
fn breakdown_uses_protocol_categories_and_authoritative_total() {
|
||||
let runtime = CanonicalMessage::text(
|
||||
"runtime",
|
||||
Role::User,
|
||||
Origin::Runtime,
|
||||
"before<rules><user_rule>r</user_rule></rules><agent_skills>s</agent_skills><subagents>a</subagents><mcp_meta_tools>m</mcp_meta_tools>after",
|
||||
);
|
||||
let snapshot = breakdown(
|
||||
1_000,
|
||||
256_000,
|
||||
None,
|
||||
"system",
|
||||
&[],
|
||||
&HashSet::new(),
|
||||
&[runtime],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.categories
|
||||
.iter()
|
||||
.map(|category| category.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
CATEGORIES
|
||||
.iter()
|
||||
.map(|category| category.0)
|
||||
.chain(std::iter::once(EASTER_EGG_CATEGORY.0))
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.categories
|
||||
.iter()
|
||||
.map(|category| category.estimated_tokens)
|
||||
.sum::<u32>(),
|
||||
1_001
|
||||
);
|
||||
for id in ["rules", "skills", "mcp", "subagents", "conversation"] {
|
||||
assert!(snapshot
|
||||
.categories
|
||||
.iter()
|
||||
.find(|category| category.id == id)
|
||||
.is_some_and(|category| category.character_count.unwrap_or(0) > 0));
|
||||
}
|
||||
assert_eq!(
|
||||
snapshot.categories[SUMMARY],
|
||||
pb::PromptTokenBreakdownCategory {
|
||||
id: "summarized_conversation".into(),
|
||||
label: "Summarized conversation".into(),
|
||||
..Default::default()
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
snapshot.categories.last().unwrap(),
|
||||
&pb::PromptTokenBreakdownCategory {
|
||||
id: "leookun".into(),
|
||||
label: "@leookun stole 1 token 😂".into(),
|
||||
estimated_tokens: 1,
|
||||
character_count: None,
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn conversation_absorbs_the_authoritative_remainder() {
|
||||
let first = breakdown(
|
||||
10_000,
|
||||
256_000,
|
||||
None,
|
||||
"system",
|
||||
&[],
|
||||
&HashSet::new(),
|
||||
&[CanonicalMessage::text(
|
||||
"user",
|
||||
Role::User,
|
||||
Origin::User,
|
||||
"short",
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
let second = breakdown(
|
||||
12_000,
|
||||
256_000,
|
||||
None,
|
||||
"system",
|
||||
&[],
|
||||
&HashSet::new(),
|
||||
&[CanonicalMessage::text(
|
||||
"user",
|
||||
Role::User,
|
||||
Origin::User,
|
||||
"a much longer conversation",
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
&first.categories[..CONVERSATION],
|
||||
&second.categories[..CONVERSATION]
|
||||
);
|
||||
assert_eq!(
|
||||
second.categories[CONVERSATION].estimated_tokens
|
||||
- first.categories[CONVERSATION].estimated_tokens,
|
||||
2_000
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
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()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
|
||||
use serde_json::json;
|
||||
use sqlx::{Connection, Row, SqliteConnection};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
const EMAIL: &str = "cursor@ai.com";
|
||||
const SIGN_UP_TYPE: &str = "Google";
|
||||
const SUBJECT: &str = "cursor-local-user";
|
||||
const MEMBERSHIP_TYPE: &str = "ultra";
|
||||
const SUBSCRIPTION_STATUS: &str = "active";
|
||||
|
||||
pub async fn inject_if_missing() -> Result<()> {
|
||||
inject_if_missing_at(&state_db_path()?).await
|
||||
}
|
||||
|
||||
fn state_db_path() -> Result<PathBuf> {
|
||||
let home = dirs::home_dir()
|
||||
.ok_or_else(|| Error::Config("cannot resolve user home directory".into()))?;
|
||||
match std::env::consts::OS {
|
||||
"macos" => {
|
||||
Ok(home.join("Library/Application Support/Cursor/User/globalStorage/state.vscdb"))
|
||||
}
|
||||
"windows" => Ok(std::env::var_os("APPDATA")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| home.join("AppData/Roaming"))
|
||||
.join("Cursor/User/globalStorage/state.vscdb")),
|
||||
"linux" => Ok(std::env::var_os("XDG_CONFIG_HOME")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| home.join(".config"))
|
||||
.join("Cursor/User/globalStorage/state.vscdb")),
|
||||
platform => Err(Error::Config(format!(
|
||||
"Cursor account injection is unsupported on {platform}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn inject_if_missing_at(path: &Path) -> Result<()> {
|
||||
if let Some(parent) = path.parent() {
|
||||
tokio::fs::create_dir_all(parent).await?;
|
||||
}
|
||||
let options = sqlx::sqlite::SqliteConnectOptions::new()
|
||||
.filename(path)
|
||||
.create_if_missing(true);
|
||||
let mut connection = SqliteConnection::connect_with(&options).await?;
|
||||
sqlx::query(
|
||||
"CREATE TABLE IF NOT EXISTS ItemTable (key TEXT UNIQUE ON CONFLICT REPLACE, value BLOB)",
|
||||
)
|
||||
.execute(&mut connection)
|
||||
.await?;
|
||||
|
||||
let account = sqlx::query("SELECT CAST(value AS TEXT) AS value FROM ItemTable WHERE key = ?")
|
||||
.bind("cursorAuth/accessToken")
|
||||
.fetch_optional(&mut connection)
|
||||
.await?;
|
||||
if account.is_some_and(|row| {
|
||||
row.try_get::<String, _>("value")
|
||||
.is_ok_and(|value| !value.trim().is_empty())
|
||||
}) {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let token = local_token()?;
|
||||
let values = [
|
||||
("cursorAuth/accessToken", token.as_str()),
|
||||
("cursorAuth/refreshToken", token.as_str()),
|
||||
("cursorAuth/cachedEmail", EMAIL),
|
||||
("cursorAuth/cachedSignUpType", SIGN_UP_TYPE),
|
||||
("cursorAuth/stripeMembershipType", MEMBERSHIP_TYPE),
|
||||
("cursorAuth/stripeSubscriptionStatus", SUBSCRIPTION_STATUS),
|
||||
];
|
||||
let mut transaction = connection.begin().await?;
|
||||
for (key, value) in values {
|
||||
sqlx::query("INSERT OR REPLACE INTO ItemTable(key, value) VALUES(?, ?)")
|
||||
.bind(key)
|
||||
.bind(value)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
}
|
||||
transaction.commit().await?;
|
||||
tracing::info!(
|
||||
email = EMAIL,
|
||||
subject = SUBJECT,
|
||||
"injected local Cursor account"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn local_token() -> Result<String> {
|
||||
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"HS256","typ":"JWT"}"#);
|
||||
let payload = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&json!({
|
||||
"sub": SUBJECT,
|
||||
"email": EMAIL,
|
||||
"type": "session",
|
||||
"iss": "cursor-client",
|
||||
"scope": "openid profile email",
|
||||
"exp": 4070908800_u64
|
||||
}))?);
|
||||
Ok(format!("{header}.{payload}.{SUBJECT}"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn injects_the_local_account_only_when_missing() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let path = directory.path().join("state.vscdb");
|
||||
|
||||
inject_if_missing_at(&path).await.unwrap();
|
||||
|
||||
let options = sqlx::sqlite::SqliteConnectOptions::new().filename(&path);
|
||||
let mut connection = SqliteConnection::connect_with(&options).await.unwrap();
|
||||
let values = sqlx::query("SELECT key, CAST(value AS TEXT) AS value FROM ItemTable")
|
||||
.fetch_all(&mut connection)
|
||||
.await
|
||||
.unwrap()
|
||||
.into_iter()
|
||||
.map(|row| (row.get::<String, _>("key"), row.get::<String, _>("value")))
|
||||
.collect::<std::collections::HashMap<_, _>>();
|
||||
let token = &values["cursorAuth/accessToken"];
|
||||
assert_eq!(values["cursorAuth/refreshToken"], *token);
|
||||
assert_eq!(values["cursorAuth/cachedEmail"], EMAIL);
|
||||
assert_eq!(values["cursorAuth/cachedSignUpType"], SIGN_UP_TYPE);
|
||||
assert_eq!(values["cursorAuth/stripeMembershipType"], MEMBERSHIP_TYPE);
|
||||
assert_eq!(
|
||||
values["cursorAuth/stripeSubscriptionStatus"],
|
||||
SUBSCRIPTION_STATUS
|
||||
);
|
||||
let payload = token.split('.').nth(1).unwrap();
|
||||
let payload: serde_json::Value =
|
||||
serde_json::from_slice(&URL_SAFE_NO_PAD.decode(payload).unwrap()).unwrap();
|
||||
assert_eq!(payload["sub"], SUBJECT);
|
||||
assert_eq!(payload["email"], EMAIL);
|
||||
assert_eq!(payload["exp"], 4070908800_u64);
|
||||
|
||||
sqlx::query("UPDATE ItemTable SET value = 'existing-token' WHERE key = ?")
|
||||
.bind("cursorAuth/accessToken")
|
||||
.execute(&mut connection)
|
||||
.await
|
||||
.unwrap();
|
||||
drop(connection);
|
||||
inject_if_missing_at(&path).await.unwrap();
|
||||
|
||||
let options = sqlx::sqlite::SqliteConnectOptions::new().filename(&path);
|
||||
let mut connection = SqliteConnection::connect_with(&options).await.unwrap();
|
||||
let token: String =
|
||||
sqlx::query_scalar("SELECT CAST(value AS TEXT) FROM ItemTable WHERE key = ?")
|
||||
.bind("cursorAuth/accessToken")
|
||||
.fetch_one(&mut connection)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(token, "existing-token");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,241 @@
|
||||
use std::{fs, path::PathBuf};
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
use std::process::Command;
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
mod windows;
|
||||
|
||||
#[cfg(unix)]
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
|
||||
use rcgen::{
|
||||
BasicConstraints, CertificateParams, DistinguishedName, DnType, IsCa, Issuer, KeyPair,
|
||||
KeyUsagePurpose, RsaKeySize, PKCS_RSA_SHA256,
|
||||
};
|
||||
#[cfg(target_os = "macos")]
|
||||
use sha1::{Digest, Sha1};
|
||||
use time::{Duration, OffsetDateTime};
|
||||
use x509_parser::prelude::FromDer;
|
||||
|
||||
use crate::{config::managed_data_dir, Error, Result};
|
||||
|
||||
use super::CaState;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CaManager {
|
||||
dir: PathBuf,
|
||||
}
|
||||
|
||||
pub struct LoadedCa {
|
||||
pub issuer: Issuer<'static, KeyPair>,
|
||||
}
|
||||
|
||||
impl CaManager {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Ok(Self {
|
||||
dir: managed_data_dir()?.join("ca"),
|
||||
})
|
||||
}
|
||||
|
||||
fn cert_path(&self) -> PathBuf {
|
||||
self.dir.join("ca.crt")
|
||||
}
|
||||
fn key_path(&self) -> PathBuf {
|
||||
self.dir.join("ca.key")
|
||||
}
|
||||
|
||||
pub fn state(&self) -> Result<CaState> {
|
||||
if !matches!(std::env::consts::OS, "macos" | "windows") {
|
||||
return Ok(CaState::Unsupported);
|
||||
}
|
||||
let cert = fs::read_to_string(self.cert_path());
|
||||
let key = fs::read_to_string(self.key_path());
|
||||
match (cert, key) {
|
||||
(Err(cert_error), Err(key_error))
|
||||
if cert_error.kind() == std::io::ErrorKind::NotFound
|
||||
&& key_error.kind() == std::io::ErrorKind::NotFound =>
|
||||
{
|
||||
Ok(CaState::Missing)
|
||||
}
|
||||
(Ok(cert), Ok(key)) => {
|
||||
if parse_issuer(&cert, &key).is_err() {
|
||||
return Ok(CaState::Invalid);
|
||||
}
|
||||
Ok(if is_installed(&cert)? {
|
||||
CaState::Ready
|
||||
} else {
|
||||
CaState::Untrusted
|
||||
})
|
||||
}
|
||||
_ => Ok(CaState::Invalid),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn load(&self) -> Result<LoadedCa> {
|
||||
let cert = fs::read_to_string(self.cert_path())?;
|
||||
let key = fs::read_to_string(self.key_path())?;
|
||||
Ok(LoadedCa {
|
||||
issuer: parse_issuer(&cert, &key)?,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn install_command(&self) -> Option<String> {
|
||||
let path = self.cert_path().to_string_lossy().replace('\'', "'\\''");
|
||||
match std::env::consts::OS {
|
||||
"macos" => dirs::home_dir().map(|_| {
|
||||
format!(
|
||||
"sudo security add-trusted-cert -d -r trustRoot -p ssl -k /Library/Keychains/System.keychain '{}'",
|
||||
path
|
||||
)
|
||||
}),
|
||||
"windows" => Some(format!(
|
||||
"certutil -addstore -f Root \"{}\"",
|
||||
self.cert_path().display()
|
||||
)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn initialize_local(&self) -> Result<()> {
|
||||
match self.state()? {
|
||||
CaState::Unsupported => {
|
||||
return Err(Error::Config(format!(
|
||||
"CA installation is not supported on {}",
|
||||
std::env::consts::OS
|
||||
)))
|
||||
}
|
||||
CaState::Invalid => {
|
||||
return Err(Error::Config("CA files are incomplete or invalid".into()))
|
||||
}
|
||||
CaState::Ready => return Ok(()),
|
||||
CaState::Missing => self.generate()?,
|
||||
CaState::Untrusted => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn generate(&self) -> Result<()> {
|
||||
fs::create_dir_all(&self.dir)?;
|
||||
#[cfg(unix)]
|
||||
fs::set_permissions(&self.dir, fs::Permissions::from_mode(0o700))?;
|
||||
|
||||
let key = KeyPair::generate_rsa_for(&PKCS_RSA_SHA256, RsaKeySize::_3072)
|
||||
.map_err(|error| Error::Config(format!("generate CA key: {error}")))?;
|
||||
let mut params = CertificateParams::new(Vec::<String>::new())
|
||||
.map_err(|error| Error::Config(format!("create CA parameters: {error}")))?;
|
||||
let mut name = DistinguishedName::new();
|
||||
name.push(DnType::CommonName, "Cursor BYOK Local CA");
|
||||
name.push(DnType::OrganizationName, "Cursor BYOK");
|
||||
params.distinguished_name = name;
|
||||
params.is_ca = IsCa::Ca(BasicConstraints::Constrained(0));
|
||||
params.key_usages = vec![
|
||||
KeyUsagePurpose::DigitalSignature,
|
||||
KeyUsagePurpose::KeyCertSign,
|
||||
KeyUsagePurpose::CrlSign,
|
||||
];
|
||||
params.not_before = OffsetDateTime::now_utc() - Duration::minutes(5);
|
||||
params.not_after = OffsetDateTime::now_utc() + Duration::days(3652);
|
||||
let cert = params
|
||||
.self_signed(&key)
|
||||
.map_err(|error| Error::Config(format!("generate CA certificate: {error}")))?;
|
||||
write_atomic(&self.key_path(), key.serialize_pem().as_bytes(), 0o600)?;
|
||||
write_atomic(&self.cert_path(), cert.pem().as_bytes(), 0o644)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_issuer(cert: &str, key: &str) -> Result<Issuer<'static, KeyPair>> {
|
||||
let key =
|
||||
KeyPair::from_pem(key).map_err(|error| Error::Config(format!("parse CA key: {error}")))?;
|
||||
let pem = pem::parse(cert).map_err(|error| Error::Config(format!("parse CA PEM: {error}")))?;
|
||||
let (_, parsed) = x509_parser::certificate::X509Certificate::from_der(pem.contents())
|
||||
.map_err(|error| Error::Config(format!("parse CA X.509 certificate: {error}")))?;
|
||||
if parsed.public_key().subject_public_key.data.as_ref() != key.public_key_raw() {
|
||||
return Err(Error::Config(
|
||||
"CA certificate and private key do not match".into(),
|
||||
));
|
||||
}
|
||||
if !parsed.validity().is_valid() {
|
||||
return Err(Error::Config(
|
||||
"CA certificate is outside its validity period".into(),
|
||||
));
|
||||
}
|
||||
if !parsed
|
||||
.basic_constraints()
|
||||
.map_err(|error| Error::Config(format!("read CA constraints: {error}")))?
|
||||
.is_some_and(|constraints| constraints.value.ca)
|
||||
{
|
||||
return Err(Error::Config("certificate is not a CA".into()));
|
||||
}
|
||||
Issuer::from_ca_cert_pem(cert, key)
|
||||
.map_err(|error| Error::Config(format!("parse CA certificate: {error}")))
|
||||
}
|
||||
|
||||
fn write_atomic(path: &std::path::Path, data: &[u8], _mode: u32) -> Result<()> {
|
||||
let temp = path.with_extension("tmp");
|
||||
fs::write(&temp, data)?;
|
||||
#[cfg(unix)]
|
||||
fs::set_permissions(&temp, fs::Permissions::from_mode(_mode))?;
|
||||
fs::rename(&temp, path)?;
|
||||
#[cfg(unix)]
|
||||
fs::set_permissions(path, fs::Permissions::from_mode(_mode))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn fingerprint(cert: &str) -> Result<String> {
|
||||
let pem = pem::parse(cert).map_err(|error| Error::Config(format!("parse CA PEM: {error}")))?;
|
||||
Ok(hex::encode_upper(Sha1::digest(pem.contents())))
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn is_installed(cert: &str) -> Result<bool> {
|
||||
let fingerprint = fingerprint(cert)?;
|
||||
for keychain in ["login.keychain-db", "/Library/Keychains/System.keychain"] {
|
||||
let output = Command::new("security")
|
||||
.args(["find-certificate", "-a", "-Z", keychain])
|
||||
.output()?;
|
||||
if output.status.success() && String::from_utf8_lossy(&output.stdout).contains(&fingerprint)
|
||||
{
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
fn is_installed(cert: &str) -> Result<bool> {
|
||||
windows::is_installed(cert)
|
||||
}
|
||||
|
||||
#[cfg(not(any(target_os = "macos", target_os = "windows")))]
|
||||
fn is_installed(_cert: &str) -> Result<bool> {
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn generated_ca_is_loadable_and_uses_private_permissions() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let manager = CaManager {
|
||||
dir: directory.path().join("ca"),
|
||||
};
|
||||
manager.generate().unwrap();
|
||||
manager.load().unwrap();
|
||||
assert!(manager.cert_path().is_file());
|
||||
assert!(manager.key_path().is_file());
|
||||
#[cfg(unix)]
|
||||
assert_eq!(
|
||||
fs::metadata(manager.key_path())
|
||||
.unwrap()
|
||||
.permissions()
|
||||
.mode()
|
||||
& 0o777,
|
||||
0o600
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
//! Native Windows system root-store access without external command-line tools.
|
||||
|
||||
use std::{ffi::c_void, io, ptr, slice};
|
||||
|
||||
use windows_sys::Win32::Security::Cryptography::{
|
||||
CertCloseStore, CertEnumCertificatesInStore, CertOpenStore, CERT_STORE_OPEN_EXISTING_FLAG,
|
||||
CERT_STORE_PROV_SYSTEM_W, CERT_STORE_READONLY_FLAG, CERT_SYSTEM_STORE_LOCAL_MACHINE,
|
||||
};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
const ROOT_STORE: [u16; 5] = [b'R' as u16, b'O' as u16, b'O' as u16, b'T' as u16, 0];
|
||||
|
||||
pub(super) fn is_installed(cert: &str) -> Result<bool> {
|
||||
let der = certificate_der(cert)?;
|
||||
let store = open_root_store()?;
|
||||
let mut context = ptr::null();
|
||||
let mut found = false;
|
||||
loop {
|
||||
context = unsafe { CertEnumCertificatesInStore(store, context) };
|
||||
if context.is_null() {
|
||||
break;
|
||||
}
|
||||
let encoded = unsafe {
|
||||
slice::from_raw_parts((*context).pbCertEncoded, (*context).cbCertEncoded as usize)
|
||||
};
|
||||
if encoded == der {
|
||||
found = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if !context.is_null() {
|
||||
unsafe { windows_sys::Win32::Security::Cryptography::CertFreeCertificateContext(context) };
|
||||
}
|
||||
close_store(store)?;
|
||||
Ok(found)
|
||||
}
|
||||
|
||||
fn certificate_der(cert: &str) -> Result<Vec<u8>> {
|
||||
pem::parse(cert)
|
||||
.map(|pem| pem.into_contents())
|
||||
.map_err(|error| Error::Config(format!("parse CA PEM: {error}")))
|
||||
}
|
||||
|
||||
fn open_root_store() -> Result<*mut c_void> {
|
||||
let flags =
|
||||
CERT_SYSTEM_STORE_LOCAL_MACHINE | CERT_STORE_OPEN_EXISTING_FLAG | CERT_STORE_READONLY_FLAG;
|
||||
let store = unsafe {
|
||||
CertOpenStore(
|
||||
CERT_STORE_PROV_SYSTEM_W,
|
||||
0,
|
||||
0,
|
||||
flags,
|
||||
ROOT_STORE.as_ptr().cast(),
|
||||
)
|
||||
};
|
||||
if store.is_null() {
|
||||
return Err(Error::Config(format!(
|
||||
"open Windows LocalMachine Root store: {}",
|
||||
io::Error::last_os_error()
|
||||
)));
|
||||
}
|
||||
Ok(store)
|
||||
}
|
||||
|
||||
fn close_store(store: *mut c_void) -> Result<()> {
|
||||
if unsafe { CertCloseStore(store, 0) } == 0 {
|
||||
return Err(Error::Config(format!(
|
||||
"close Windows certificate store: {}",
|
||||
io::Error::last_os_error()
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
mod account;
|
||||
mod ca;
|
||||
mod proxy;
|
||||
mod settings;
|
||||
|
||||
use std::{net::SocketAddr, sync::Arc};
|
||||
|
||||
use parking_lot::RwLock;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::{store::Store, Error, Result};
|
||||
|
||||
use self::{ca::CaManager, proxy::ProxyRuntime};
|
||||
|
||||
pub(crate) fn proxy_host_allowed(host: &str) -> bool {
|
||||
proxy::is_cursor_host(host)
|
||||
}
|
||||
|
||||
fn integration_prerequisites_ready(ca: &CaState, backend_ready: bool) -> bool {
|
||||
matches!(ca, CaState::Ready) && backend_ready
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum CaState {
|
||||
Missing,
|
||||
Untrusted,
|
||||
Ready,
|
||||
Invalid,
|
||||
Unsupported,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum IntegrationState {
|
||||
Disabled,
|
||||
Enabled,
|
||||
Degraded,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct CursorHarnessStatus {
|
||||
pub platform: &'static str,
|
||||
pub ca: CaState,
|
||||
pub configured_models: usize,
|
||||
pub enabled_models: usize,
|
||||
pub integration: IntegrationState,
|
||||
pub proxy_url: Option<String>,
|
||||
pub ca_install_command: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize)]
|
||||
pub struct SetEnabled {
|
||||
pub enabled: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorHarness {
|
||||
inner: Arc<Inner>,
|
||||
}
|
||||
|
||||
struct Inner {
|
||||
store: Store,
|
||||
ca: CaManager,
|
||||
ca_initialization: Mutex<()>,
|
||||
backend_addr: RwLock<Option<SocketAddr>>,
|
||||
proxy: Mutex<ProxyRuntime>,
|
||||
}
|
||||
|
||||
impl CursorHarness {
|
||||
pub fn new(store: Store) -> Result<Self> {
|
||||
Ok(Self {
|
||||
inner: Arc::new(Inner {
|
||||
store,
|
||||
ca: CaManager::managed()?,
|
||||
ca_initialization: Mutex::new(()),
|
||||
backend_addr: RwLock::new(None),
|
||||
proxy: Mutex::new(ProxyRuntime::default()),
|
||||
}),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn set_backend_addr(&self, addr: SocketAddr) {
|
||||
*self.inner.backend_addr.write() = Some(addr);
|
||||
}
|
||||
|
||||
pub async fn cleanup_stale_settings(&self) -> Result<()> {
|
||||
settings::clear_stale_managed_settings()
|
||||
}
|
||||
|
||||
pub async fn status(&self) -> Result<CursorHarnessStatus> {
|
||||
let models = self.inner.store.provider_models(false).await?;
|
||||
let configured_models = models.len();
|
||||
let enabled_models = models.iter().filter(|model| model.enabled).count();
|
||||
let ca = self.inner.ca.state()?;
|
||||
if integration_prerequisites_ready(&ca, self.inner.backend_addr.read().is_some()) {
|
||||
self.enable().await?;
|
||||
}
|
||||
let proxy = self.inner.proxy.lock().await;
|
||||
let proxy_url = proxy.url();
|
||||
let settings_applied = proxy_url
|
||||
.as_deref()
|
||||
.map(settings::settings_match)
|
||||
.transpose()?
|
||||
.unwrap_or(false);
|
||||
let integration = match (proxy.running(), settings_applied) {
|
||||
(false, false) => IntegrationState::Disabled,
|
||||
(true, true) => IntegrationState::Enabled,
|
||||
_ => IntegrationState::Degraded,
|
||||
};
|
||||
Ok(CursorHarnessStatus {
|
||||
platform: std::env::consts::OS,
|
||||
ca,
|
||||
configured_models,
|
||||
enabled_models,
|
||||
integration,
|
||||
proxy_url,
|
||||
ca_install_command: self.inner.ca.install_command(),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn initialize_ca(&self) -> Result<CursorHarnessStatus> {
|
||||
let _initialization = self.inner.ca_initialization.lock().await;
|
||||
let manager = self.inner.ca.clone();
|
||||
tokio::task::spawn_blocking(move || manager.initialize_local())
|
||||
.await
|
||||
.map_err(|error| Error::Store(format!("CA initialization task failed: {error}")))??;
|
||||
self.status().await
|
||||
}
|
||||
|
||||
pub async fn set_enabled(&self, enabled: bool) -> Result<CursorHarnessStatus> {
|
||||
if enabled {
|
||||
self.enable().await?;
|
||||
} else {
|
||||
self.disable().await?;
|
||||
}
|
||||
self.status().await
|
||||
}
|
||||
|
||||
async fn enable(&self) -> Result<()> {
|
||||
if !matches!(self.inner.ca.state()?, CaState::Ready) {
|
||||
return Err(Error::Config(
|
||||
"initialize and trust the CA before enabling Cursor".into(),
|
||||
));
|
||||
}
|
||||
let backend_addr = self
|
||||
.inner
|
||||
.backend_addr
|
||||
.read()
|
||||
.ok_or_else(|| Error::Config("desktop management server is not ready".into()))?;
|
||||
let mut proxy = self.inner.proxy.lock().await;
|
||||
if proxy.running() {
|
||||
if let Some(url) = proxy.url() {
|
||||
apply_cursor_configuration(&url).await?;
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
let ca = self.inner.ca.load()?;
|
||||
let requested_port = self.inner.store.port_settings().await?.proxy_port;
|
||||
let (url, actual_port) = proxy.start(backend_addr, ca, requested_port).await?;
|
||||
if let Err(error) = self.inner.store.set_proxy_port(actual_port).await {
|
||||
proxy.stop().await;
|
||||
return Err(error);
|
||||
}
|
||||
if let Err(error) = apply_cursor_configuration(&url).await {
|
||||
proxy.stop().await;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn disable(&self) -> Result<()> {
|
||||
settings::clear_proxy_settings()?;
|
||||
self.inner.proxy.lock().await.stop().await;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn apply_cursor_configuration(proxy_url: &str) -> Result<()> {
|
||||
account::inject_if_missing().await?;
|
||||
settings::write_proxy_settings(proxy_url)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn automatic_integration_requires_only_a_ready_ca_and_backend() {
|
||||
assert!(integration_prerequisites_ready(&CaState::Ready, true));
|
||||
assert!(!integration_prerequisites_ready(&CaState::Ready, false));
|
||||
assert!(!integration_prerequisites_ready(&CaState::Missing, true));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use hudsucker::{
|
||||
certificate_authority::RcgenAuthority,
|
||||
hyper::{Request, Uri},
|
||||
rustls::crypto::aws_lc_rs,
|
||||
Body, HttpContext, HttpHandler, Proxy, RequestOrResponse,
|
||||
};
|
||||
use tokio::{net::TcpListener, sync::oneshot, task::JoinHandle};
|
||||
|
||||
use crate::{cursor::proxy::UPSTREAM_URL_HEADER, Error, Result};
|
||||
|
||||
use super::ca::LoadedCa;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct ProxyRuntime {
|
||||
url: Option<String>,
|
||||
port: Option<u16>,
|
||||
stop: Option<oneshot::Sender<()>>,
|
||||
task: Option<JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl ProxyRuntime {
|
||||
pub fn running(&self) -> bool {
|
||||
self.task.as_ref().is_some_and(|task| !task.is_finished())
|
||||
}
|
||||
pub fn url(&self) -> Option<String> {
|
||||
self.running().then(|| self.url.clone()).flatten()
|
||||
}
|
||||
|
||||
pub async fn start(
|
||||
&mut self,
|
||||
backend: SocketAddr,
|
||||
ca: LoadedCa,
|
||||
requested_port: u16,
|
||||
) -> Result<(String, u16)> {
|
||||
if let Some(url) = self.url() {
|
||||
return Ok((url, self.port.unwrap_or_default()));
|
||||
}
|
||||
let listener = bind_proxy_listener(requested_port).await?;
|
||||
let address = listener.local_addr()?;
|
||||
let (stop, done) = oneshot::channel();
|
||||
let authority = RcgenAuthority::new(ca.issuer, 1_000, aws_lc_rs::default_provider());
|
||||
let proxy = Proxy::builder()
|
||||
.with_listener(listener)
|
||||
.with_ca(authority)
|
||||
.with_rustls_connector(aws_lc_rs::default_provider())
|
||||
.with_http_handler(CursorRelay { backend })
|
||||
.with_graceful_shutdown(async move {
|
||||
let _ = done.await;
|
||||
})
|
||||
.build()
|
||||
.map_err(|error| Error::Store(format!("build Cursor proxy: {error}")))?;
|
||||
self.stop = Some(stop);
|
||||
self.url = Some(format!("http://{address}"));
|
||||
self.port = Some(address.port());
|
||||
self.task = Some(tokio::spawn(async move {
|
||||
if let Err(error) = proxy.start().await {
|
||||
tracing::error!(%error, "Cursor proxy stopped unexpectedly");
|
||||
}
|
||||
}));
|
||||
Ok((self.url.clone().unwrap(), address.port()))
|
||||
}
|
||||
|
||||
pub async fn stop(&mut self) {
|
||||
if let Some(stop) = self.stop.take() {
|
||||
let _ = stop.send(());
|
||||
}
|
||||
if let Some(task) = self.task.take() {
|
||||
let _ = tokio::time::timeout(std::time::Duration::from_secs(5), task).await;
|
||||
}
|
||||
self.url = None;
|
||||
self.port = None;
|
||||
}
|
||||
}
|
||||
|
||||
async fn bind_proxy_listener(requested_port: u16) -> Result<TcpListener> {
|
||||
let requested = SocketAddr::from(([127, 0, 0, 1], requested_port));
|
||||
match TcpListener::bind(requested).await {
|
||||
Ok(listener) => Ok(listener),
|
||||
Err(error) if requested_port != 0 => {
|
||||
tracing::warn!(%requested, %error, "configured proxy port unavailable; selecting a random port");
|
||||
Ok(TcpListener::bind("127.0.0.1:0").await?)
|
||||
}
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct CursorRelay {
|
||||
backend: SocketAddr,
|
||||
}
|
||||
|
||||
impl HttpHandler for CursorRelay {
|
||||
async fn handle_request(
|
||||
&mut self,
|
||||
_ctx: &HttpContext,
|
||||
mut request: Request<Body>,
|
||||
) -> RequestOrResponse {
|
||||
let original = request.uri().clone();
|
||||
if is_cursor_host(original.host().unwrap_or_default()) && is_local_path(original.path()) {
|
||||
if let Ok(value) = original.to_string().parse() {
|
||||
request.headers_mut().insert(UPSTREAM_URL_HEADER, value);
|
||||
}
|
||||
let path = original
|
||||
.path_and_query()
|
||||
.map(|value| value.as_str())
|
||||
.unwrap_or("/");
|
||||
if let Ok(uri) = format!("http://{}{}", self.backend, path).parse::<Uri>() {
|
||||
*request.uri_mut() = uri;
|
||||
}
|
||||
}
|
||||
request.into()
|
||||
}
|
||||
|
||||
async fn should_intercept_connect(
|
||||
&mut self,
|
||||
_ctx: &HttpContext,
|
||||
request: &Request<Body>,
|
||||
) -> bool {
|
||||
request
|
||||
.uri()
|
||||
.authority()
|
||||
.is_some_and(|authority| is_cursor_host(authority.host()))
|
||||
}
|
||||
|
||||
async fn should_intercept_tls(
|
||||
&mut self,
|
||||
_ctx: &HttpContext,
|
||||
hello: hudsucker::rustls::server::ClientHello<'_>,
|
||||
) -> bool {
|
||||
hello.server_name().is_some_and(is_cursor_host)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_cursor_host(host: &str) -> bool {
|
||||
let host = host.trim_end_matches('.').to_ascii_lowercase();
|
||||
matches!(host.as_str(), "api2.cursor.sh" | "api3.cursor.sh") || host.ends_with(".cursor.sh")
|
||||
}
|
||||
|
||||
fn is_local_path(path: &str) -> bool {
|
||||
matches!(
|
||||
path,
|
||||
"/agent.v1.AgentService/RunSSE"
|
||||
| "/aiserver.v1.BidiService/BidiAppend"
|
||||
| "/aiserver.v1.AiService/AvailableModels"
|
||||
| "/agent.v1.AgentService/GetUsableModels"
|
||||
| "/aiserver.v1.AiService/GetUsableModels"
|
||||
| "/aiserver.v1.AuthService/GetEmail"
|
||||
| "/aiserver.v1.DashboardService/GetMe"
|
||||
| "/aiserver.v1.DashboardService/GetTeams"
|
||||
| "/aiserver.v1.DashboardService/GetUserProfile"
|
||||
| "/aiserver.v1.DashboardService/GetCurrentPeriodUsage"
|
||||
| "/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants"
|
||||
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
|
||||
| "/auth/full_stripe_profile"
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn proxy_listener_falls_back_when_configured_port_is_busy() {
|
||||
let occupied = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let requested_port = occupied.local_addr().unwrap().port();
|
||||
let listener = bind_proxy_listener(requested_port).await.unwrap();
|
||||
assert_ne!(listener.local_addr().unwrap().port(), requested_port);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn limits_interception_to_cursor_hosts_and_local_paths() {
|
||||
assert!(is_cursor_host("api2.cursor.sh"));
|
||||
assert!(is_cursor_host("repo42.cursor.sh"));
|
||||
assert!(!is_cursor_host("example.com"));
|
||||
assert!(is_local_path("/agent.v1.AgentService/RunSSE"));
|
||||
assert!(is_local_path(
|
||||
"/aiserver.v1.AnalyticsService/BootstrapStatsig"
|
||||
));
|
||||
assert!(!is_local_path("/unrelated"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
use std::{collections::BTreeMap, fs, path::PathBuf};
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
const KEYS: [&str; 5] = [
|
||||
"http.proxy",
|
||||
"http.proxyKerberosServicePrincipal",
|
||||
"http.proxySupport",
|
||||
"cursor.general.disableHttp2",
|
||||
"http.experimental.systemCertificatesV2",
|
||||
];
|
||||
|
||||
fn path() -> Result<PathBuf> {
|
||||
let home = dirs::home_dir()
|
||||
.ok_or_else(|| Error::Config("cannot resolve user home directory".into()))?;
|
||||
match std::env::consts::OS {
|
||||
"macos" => Ok(home.join("Library/Application Support/Cursor/User/settings.json")),
|
||||
"windows" => Ok(std::env::var_os("APPDATA")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| home.join("AppData/Roaming"))
|
||||
.join("Cursor/User/settings.json")),
|
||||
"linux" => Ok(std::env::var_os("XDG_CONFIG_HOME")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| home.join(".config"))
|
||||
.join("Cursor/User/settings.json")),
|
||||
platform => Err(Error::Config(format!(
|
||||
"Cursor settings are unsupported on {platform}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn read() -> Result<BTreeMap<String, Value>> {
|
||||
let path = path()?;
|
||||
let data = match fs::read_to_string(path) {
|
||||
Ok(data) => data,
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(BTreeMap::new()),
|
||||
Err(error) => return Err(error.into()),
|
||||
};
|
||||
if data.trim().is_empty() {
|
||||
return Ok(BTreeMap::new());
|
||||
}
|
||||
json5::from_str(&data)
|
||||
.map_err(|error| Error::Config(format!("parse Cursor settings JSONC: {error}")))
|
||||
}
|
||||
|
||||
fn write(settings: &BTreeMap<String, Value>) -> Result<()> {
|
||||
let path = path()?;
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent)?;
|
||||
}
|
||||
let data = serde_json::to_vec_pretty(settings)?;
|
||||
let temp = path.with_extension("json.tmp");
|
||||
fs::write(&temp, [data.as_slice(), b"\n"].concat())?;
|
||||
fs::rename(temp, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn write_proxy_settings(proxy_url: &str) -> Result<()> {
|
||||
let mut settings = read()?;
|
||||
settings.insert(KEYS[0].into(), Value::String(proxy_url.into()));
|
||||
settings.insert(KEYS[1].into(), Value::String(proxy_url.into()));
|
||||
settings.insert(KEYS[2].into(), Value::String("on".into()));
|
||||
settings.insert(KEYS[3].into(), Value::Bool(true));
|
||||
settings.insert(KEYS[4].into(), Value::Bool(true));
|
||||
write(&settings)
|
||||
}
|
||||
|
||||
pub fn clear_proxy_settings() -> Result<()> {
|
||||
let mut settings = read()?;
|
||||
let before = settings.len();
|
||||
for key in KEYS {
|
||||
settings.remove(key);
|
||||
}
|
||||
if settings.len() != before {
|
||||
write(&settings)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn settings_match(proxy_url: &str) -> Result<bool> {
|
||||
let settings = read()?;
|
||||
Ok(
|
||||
settings.get(KEYS[0]) == Some(&Value::String(proxy_url.into()))
|
||||
&& settings.get(KEYS[1]) == Some(&Value::String(proxy_url.into()))
|
||||
&& settings.get(KEYS[2]) == Some(&Value::String("on".into()))
|
||||
&& settings.get(KEYS[3]) == Some(&Value::Bool(true))
|
||||
&& settings.get(KEYS[4]) == Some(&Value::Bool(true)),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn clear_stale_managed_settings() -> Result<()> {
|
||||
let settings = read()?;
|
||||
let managed_signature = settings.get(KEYS[2]) == Some(&Value::String("on".into()))
|
||||
&& settings.get(KEYS[3]) == Some(&Value::Bool(true))
|
||||
&& settings.get(KEYS[4]) == Some(&Value::Bool(true));
|
||||
let loopback = settings
|
||||
.get(KEYS[0])
|
||||
.and_then(Value::as_str)
|
||||
.and_then(|value| value.parse::<reqwest::Url>().ok())
|
||||
.and_then(|url| url.host_str().map(str::to_owned))
|
||||
.is_some_and(|host| matches!(host.as_str(), "127.0.0.1" | "localhost" | "::1"));
|
||||
if managed_signature && loopback {
|
||||
clear_proxy_settings()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[test]
|
||||
fn json5_accepts_cursor_jsonc() {
|
||||
let parsed: std::collections::BTreeMap<String, serde_json::Value> =
|
||||
json5::from_str("{ // note\n 'a': 1, }").unwrap();
|
||||
assert_eq!(parsed["a"], 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
pub mod app;
|
||||
pub mod client;
|
||||
pub mod config;
|
||||
pub mod control;
|
||||
pub mod cursor;
|
||||
pub mod error;
|
||||
pub mod harness;
|
||||
pub mod model;
|
||||
pub mod network;
|
||||
pub mod provider;
|
||||
pub mod run;
|
||||
pub mod store;
|
||||
pub mod web;
|
||||
|
||||
pub use app::App;
|
||||
pub use config::Config;
|
||||
pub use error::{Error, Result};
|
||||
@@ -0,0 +1,60 @@
|
||||
use std::fmt;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
macro_rules! string_id {
|
||||
($name:ident) => {
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord)]
|
||||
#[serde(transparent)]
|
||||
pub struct $name(pub String);
|
||||
|
||||
impl $name {
|
||||
pub fn new(value: impl Into<String>) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for $name {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for $name {
|
||||
fn from(value: String) -> Self {
|
||||
Self(value)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for $name {
|
||||
fn from(value: &str) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
string_id!(ConversationId);
|
||||
string_id!(RunId);
|
||||
string_id!(ToolRoundId);
|
||||
|
||||
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord)]
|
||||
#[serde(transparent)]
|
||||
pub struct RevisionId(pub i64);
|
||||
|
||||
impl fmt::Display for RevisionId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct Conversation {
|
||||
pub conversation_id: ConversationId,
|
||||
pub current_revision_id: RevisionId,
|
||||
pub active_run_id: Option<RunId>,
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct CursorRunTraceSummary {
|
||||
pub request_id: String,
|
||||
pub conversation_id: Option<String>,
|
||||
pub route: String,
|
||||
pub model_id: Option<String>,
|
||||
pub status: String,
|
||||
pub request_bytes: i64,
|
||||
pub response_bytes: i64,
|
||||
pub response_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub received_at_ms: i64,
|
||||
pub first_response_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub error_message: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct CursorRunTraceArtifact {
|
||||
pub seq: i64,
|
||||
pub artifact_type: String,
|
||||
pub source: String,
|
||||
pub metadata: serde_json::Value,
|
||||
pub created_at_ms: i64,
|
||||
pub data: Vec<u8>,
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{ModelSpec, ProjectedMessage, ToolDefinition};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct PromptSpec {
|
||||
pub instructions: String,
|
||||
pub tools: Vec<ToolDefinition>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ModelRequest {
|
||||
pub prompt: PromptSpec,
|
||||
pub model: ModelSpec,
|
||||
pub history: Vec<ProjectedMessage>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ModelInvocation {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: u64,
|
||||
pub request: ModelRequest,
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
use serde::Serialize;
|
||||
|
||||
use super::ProviderType;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct NewLlmCall {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub provider_url: String,
|
||||
pub request_type: ProviderType,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: bool,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallSummary {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: Option<String>,
|
||||
pub provider_type: String,
|
||||
pub provider_url: String,
|
||||
pub request_type: String,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: Option<bool>,
|
||||
pub status: String,
|
||||
pub finish_reason: Option<String>,
|
||||
pub created_at_ms: i64,
|
||||
pub request_started_at_ms: Option<i64>,
|
||||
pub response_headers_at_ms: Option<i64>,
|
||||
pub first_event_at_ms: Option<i64>,
|
||||
pub first_text_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub queue_ms: Option<i64>,
|
||||
pub ttfb_ms: Option<i64>,
|
||||
pub ttft_ms: Option<i64>,
|
||||
pub duration_ms: Option<i64>,
|
||||
pub input_tokens: Option<i64>,
|
||||
pub output_tokens: Option<i64>,
|
||||
pub total_tokens: Option<i64>,
|
||||
pub cache_read_tokens: Option<i64>,
|
||||
pub cache_write_tokens: Option<i64>,
|
||||
pub reasoning_tokens: Option<i64>,
|
||||
pub usage: Option<serde_json::Value>,
|
||||
pub message_count: i64,
|
||||
pub tool_count: i64,
|
||||
pub request_bytes: Option<i64>,
|
||||
pub response_bytes: i64,
|
||||
pub stream_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub error_kind: Option<String>,
|
||||
pub error_message: Option<String>,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallRequest {
|
||||
pub headers: serde_json::Value,
|
||||
pub body: serde_json::Value,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallResponseChunk {
|
||||
pub seq: i64,
|
||||
pub received_offset_ms: i64,
|
||||
pub data: String,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user