mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-18 03:57:06 +08:00
all tools
This commit is contained in:
@@ -0,0 +1,79 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{connect, proto::agent::v1 as pb},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
run::RunRegistry,
|
||||
store::BlobId,
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_dependency_is_not_ready_until_blob_ack() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
store.create_pending_run("request").await.unwrap();
|
||||
let id = BlobId::digest(b"tool result");
|
||||
let key = format!("blob:{}", id.to_base64());
|
||||
store
|
||||
.enqueue_outbox("request", &key, "kv_set", b"tool result", &[])
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(!store
|
||||
.dependencies_acked("request", std::slice::from_ref(&id))
|
||||
.await
|
||||
.unwrap());
|
||||
assert!(store.ack_outbox("request", &key).await.unwrap());
|
||||
assert!(store.dependencies_acked("request", &[id]).await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pending_blob_outbox_is_replayed_when_the_run_is_recreated() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
store.create_pending_run("recover-request").await.unwrap();
|
||||
let data = b"durable tool result";
|
||||
let blob_id = store.put_blob(data, &[]).await.unwrap();
|
||||
store
|
||||
.enqueue_outbox(
|
||||
"recover-request",
|
||||
&format!("blob:{}", blob_id.to_base64()),
|
||||
"kv_set",
|
||||
data,
|
||||
&[],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
"test".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("recover-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(2), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
let message = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message else {
|
||||
panic!("expected replayed KV SET")
|
||||
};
|
||||
let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = kv.message else {
|
||||
panic!("expected SetBlobArgs")
|
||||
};
|
||||
assert_eq!(set.blob_id, blob_id.as_bytes());
|
||||
assert_eq!(set.blob_data, data);
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
#[path = "support/fake_cursor.rs"]
|
||||
mod fake_cursor;
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{io::Write, sync::Arc};
|
||||
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{header, Request, StatusCode},
|
||||
};
|
||||
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect, handlers,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
run::RunRegistry,
|
||||
};
|
||||
use flate2::{write::GzEncoder, Compression};
|
||||
use prost::Message;
|
||||
use tower::ServiceExt;
|
||||
|
||||
#[test]
|
||||
fn connect_envelope_is_flag_plus_big_endian_length_plus_protobuf() {
|
||||
let message = pb::BidiRequestId {
|
||||
request_id: "abc".into(),
|
||||
};
|
||||
let frame = connect::encode_message(&message).unwrap();
|
||||
assert_eq!(&frame[..5], &[0, 0, 0, 0, 5]);
|
||||
let decoded: pb::BidiRequestId = fake_cursor::decode_single(&frame).unwrap();
|
||||
assert_eq!(decoded.request_id, "abc");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn end_stream_matches_captured_connect_shape() {
|
||||
assert_eq!(
|
||||
connect::encode_end_stream().as_ref(),
|
||||
&[2, 0, 0, 0, 2, b'{', b'}']
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_end_stream_is_flagged_json_not_protobuf() {
|
||||
let frame = connect::encode_error_end_stream(&connect::ConnectStreamError {
|
||||
code: connect::ConnectCode::Unavailable,
|
||||
message: "overloaded".into(),
|
||||
details: vec![connect::ConnectErrorDetail {
|
||||
type_name: "aiserver.v1.ErrorDetails".into(),
|
||||
value: "AQ".into(),
|
||||
}],
|
||||
})
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
assert_eq!(flags, connect::END_STREAM_FLAG);
|
||||
let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
|
||||
assert_eq!(json["error"]["code"], "unavailable");
|
||||
assert_eq!(json["error"]["message"], "overloaded");
|
||||
assert_eq!(
|
||||
json["error"]["details"][0]["type"],
|
||||
"aiserver.v1.ErrorDetails"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cursor_error_details_subset_decodes_captured_wire_value() {
|
||||
let captured = "CAISVQoUQXV0aGVudGljYXRpb24gZXJyb3ISMklmIHlvdSBhcmUgbG9nZ2VkIGluLCB0cnkgbG9nZ2luZyBvdXQgYW5kIGJhY2sgaW4uIABSBwoFbG9naW4YAQ";
|
||||
let bytes = STANDARD_NO_PAD.decode(captured).unwrap();
|
||||
let details = ai::ErrorDetails::decode(bytes.as_slice()).unwrap();
|
||||
assert_eq!(details.error, 2, "ERROR_NOT_LOGGED_IN");
|
||||
assert_eq!(details.is_expected, Some(true));
|
||||
let custom = details.details.unwrap();
|
||||
assert_eq!(custom.title, "Authentication error");
|
||||
assert_eq!(custom.is_retryable, Some(false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn captured_kv_ack_hex_decodes_as_agent_client_message() {
|
||||
let bytes = hex::decode("1a0408011a00").unwrap();
|
||||
let message = pb::AgentClientMessage::decode(bytes.as_slice()).unwrap();
|
||||
let Some(pb::agent_client_message::Message::KvClientMessage(kv)) = message.message else {
|
||||
panic!("expected KV client message")
|
||||
};
|
||||
assert_eq!(kv.id, 1);
|
||||
assert!(matches!(
|
||||
kv.message,
|
||||
Some(pb::kv_client_message::Message::SetBlobResult(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let wire = ai::BidiAppendRequest {
|
||||
request_id: Some(ai::BidiRequestId {
|
||||
request_id: "gzip-request".into(),
|
||||
}),
|
||||
..Default::default()
|
||||
}
|
||||
.encode_to_vec();
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder.write_all(&wire).unwrap();
|
||||
let compressed = encoder.finish().unwrap();
|
||||
|
||||
let response = handlers::router(registry)
|
||||
.oneshot(
|
||||
Request::post("/aiserver.v1.BidiService/BidiAppend")
|
||||
.header(header::CONTENT_TYPE, "application/proto")
|
||||
.header(header::CONTENT_ENCODING, "gzip")
|
||||
.body(Body::from(compressed))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
let text = std::str::from_utf8(&body).unwrap();
|
||||
assert!(text.contains("BidiAppend contains no AgentClientMessage"));
|
||||
assert!(!text.contains("protobuf decode error"));
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
model::{MessageContent, Role},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
provider::{FinishReason, ResponseEvent},
|
||||
run::{RunCommand, RunRegistry},
|
||||
Error,
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_failure_checkpoints_then_returns_structured_error_and_closes() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push_error(Error::Provider("provider failed".into()));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("failed-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut saw_checkpoint = false;
|
||||
let mut saw_turn_ended = false;
|
||||
let error_json = loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
break serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(_)) => {
|
||||
saw_checkpoint = true;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
if matches!(
|
||||
update.message,
|
||||
Some(pb::interaction_update::Message::TurnEnded(_))
|
||||
) {
|
||||
saw_turn_ended = true;
|
||||
}
|
||||
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
|
||||
assert!(!delta.text.contains("Cursor server error"));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
|
||||
assert!(saw_checkpoint);
|
||||
assert!(!saw_turn_ended);
|
||||
assert_eq!(error_json["error"]["code"], "unavailable");
|
||||
let detail = &error_json["error"]["details"][0];
|
||||
assert_eq!(detail["type"], "aiserver.v1.ErrorDetails");
|
||||
let encoded = detail["value"].as_str().unwrap();
|
||||
assert!(!encoded.ends_with('='));
|
||||
let decoded = STANDARD_NO_PAD.decode(encoded).unwrap();
|
||||
let decoded = ai::ErrorDetails::decode(decoded.as_slice()).unwrap();
|
||||
assert_eq!(
|
||||
decoded.error,
|
||||
ai::error_details::Error::ProviderError as i32
|
||||
);
|
||||
let custom = decoded.details.unwrap();
|
||||
assert_eq!(custom.title, "Server Error");
|
||||
assert_eq!(custom.is_retryable, Some(true));
|
||||
assert_eq!(custom.should_show_immediate_error, Some(false));
|
||||
assert_eq!(
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
|
||||
let messages = store.load_messages("failed-conversation").await.unwrap();
|
||||
assert!(messages.iter().any(|message| message.role == Role::User));
|
||||
assert!(!messages.iter().any(|message| {
|
||||
matches!(
|
||||
&message.content,
|
||||
MessageContent::Assistant { text, .. } if text.contains("Cursor server error")
|
||||
)
|
||||
}));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ResponseEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallEnd { index: 0 },
|
||||
ResponseEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("protocol-failed-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(protocol_client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut saw_turn_ended = false;
|
||||
let error_json = loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before Error EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
break serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
// An unknown numeric bridge id is a runtime protocol error.
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: exec.id + 1_000,
|
||||
exec_id: String::new(),
|
||||
message: None,
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
if matches!(
|
||||
update.message,
|
||||
Some(pb::interaction_update::Message::TurnEnded(_))
|
||||
) {
|
||||
saw_turn_ended = true;
|
||||
}
|
||||
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
|
||||
assert!(!delta.text.contains("unknown tool result"));
|
||||
assert!(!delta.text.contains("protocol error"));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
|
||||
assert!(!saw_turn_ended);
|
||||
assert_eq!(error_json["error"]["code"], "invalid_argument");
|
||||
assert_eq!(
|
||||
error_json["error"]["message"],
|
||||
"protocol error: unknown ExecClientMessage id: 1001"
|
||||
);
|
||||
assert_eq!(
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
fn client_run() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "hello".into(),
|
||||
message_id: "failed-user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("failed-conversation".into()),
|
||||
run_id: Some("failed-request".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn protocol_client_run() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "read it".into(),
|
||||
message_id: "protocol-failed-user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("protocol-failed-conversation".into()),
|
||||
run_id: Some("protocol-failed-request".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn kv_ack(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id,
|
||||
message: Some(pb::kv_client_message::Message::SetBlobResult(
|
||||
pb::SetBlobResult { error: None },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{connect, proto::agent::v1 as pb},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
provider::{FinishReason, ResponseEvent},
|
||||
run::{RunCommand, RunRegistry},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_new_revision_invalidates_late_events_from_the_old_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let first = store.begin_revision("conversation").await.unwrap();
|
||||
let second = store.begin_revision("conversation").await.unwrap();
|
||||
assert!(second > first);
|
||||
assert!(!store
|
||||
.revision_is_current("conversation", first)
|
||||
.await
|
||||
.unwrap());
|
||||
assert!(store
|
||||
.revision_is_current("conversation", second)
|
||||
.await
|
||||
.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn registry_shutdown_cancels_runs_and_closes_run_sse_outputs() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("active-run").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
|
||||
registry.shutdown().await;
|
||||
|
||||
assert!(handle.cancellation().is_cancelled());
|
||||
let terminal = output.recv().await.expect("canceled EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&terminal).unwrap().pop().unwrap();
|
||||
assert_eq!(flags, connect::END_STREAM_FLAG);
|
||||
let payload: serde_json::Value = serde_json::from_slice(&payload).unwrap();
|
||||
assert_eq!(payload["error"]["code"], "canceled");
|
||||
assert_eq!(output.recv().await, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancel_aborts_active_exec_before_canceled_end_stream() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ResponseEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallEnd { index: 0 },
|
||||
ResponseEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("cancel-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let exec_id = loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => break exec.id,
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
|
||||
handle.cancel();
|
||||
let mut saw_abort = false;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before canceled EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
|
||||
assert_eq!(json["error"]["code"], "canceled");
|
||||
assert!(saw_abort, "ExecServerAbort must precede canceled EndStream");
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) =
|
||||
server.message
|
||||
{
|
||||
let Some(pb::exec_server_control_message::Message::Abort(abort)) = control.message
|
||||
else {
|
||||
panic!("expected ExecServerAbort")
|
||||
};
|
||||
assert_eq!(abort.id, exec_id);
|
||||
saw_abort = true;
|
||||
}
|
||||
}
|
||||
assert_eq!(output.recv().await, None);
|
||||
}
|
||||
|
||||
fn client_run() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "read".into(),
|
||||
message_id: "cancel-user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("cancel-conversation".into()),
|
||||
run_id: Some("cancel-request".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn kv_ack(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id,
|
||||
message: Some(pb::kv_client_message::Message::SetBlobResult(
|
||||
pb::SetBlobResult { error: None },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use cursor_server::{
|
||||
model::{CanonicalMessage, MessageContent, Origin, Role, ToolCallContent, ToolResultContent},
|
||||
prompting::{project_messages, Mode, PromptAssets, PromptCompiler, ToolDefinition},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn projecting_an_append_only_context_preserves_the_complete_prefix() {
|
||||
let first = vec![fixtures::user("u1", "one")];
|
||||
let mut second = first.clone();
|
||||
second.push(fixtures::user("u2", "two"));
|
||||
let projected_first = project_messages(&first).unwrap();
|
||||
let projected_second = project_messages(&second).unwrap();
|
||||
assert_eq!(projected_first, projected_second[..projected_first.len()]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_tool_result_is_projected_as_string_content() {
|
||||
let object = serde_json::json!({"merge": false, "todos": []});
|
||||
let messages = vec![
|
||||
tool_result("object", object.clone()),
|
||||
tool_result("string", serde_json::Value::String("plain text".into())),
|
||||
];
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
|
||||
let object_text = projected[0]
|
||||
.content
|
||||
.as_str()
|
||||
.expect("object ToolResult must be JSON-encoded into a string");
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(object_text).unwrap(),
|
||||
object
|
||||
);
|
||||
assert_eq!(projected[1].content.as_str(), Some("plain text"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assistant_text_and_thinking_remain_separate_during_projection() {
|
||||
let messages = vec![CanonicalMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
origin: Origin::Assistant,
|
||||
content: MessageContent::Assistant {
|
||||
text: "visible answer".into(),
|
||||
thinking: "private reasoning".into(),
|
||||
model_call_id: Some("model-call".into()),
|
||||
tool_calls: Vec::new(),
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}];
|
||||
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
assert_eq!(projected[0].content.as_str(), Some("visible answer"));
|
||||
assert_eq!(projected[0].thinking.as_deref(), Some("private reasoning"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_tool_pairs_reconstruct_the_original_provider_assistant_message() {
|
||||
let messages = vec![
|
||||
assistant_tool_pair(
|
||||
"assistant-second",
|
||||
"model-call",
|
||||
1,
|
||||
"call-second",
|
||||
"visible answer",
|
||||
"complete reasoning",
|
||||
),
|
||||
tool_result_with_call("result-second", "call-second", "second"),
|
||||
assistant_tool_pair("assistant-first", "model-call", 0, "call-first", "", ""),
|
||||
tool_result_with_call("result-first", "call-first", "first"),
|
||||
];
|
||||
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
|
||||
assert_eq!(projected.len(), 3);
|
||||
assert_eq!(projected[0].role, "assistant");
|
||||
assert_eq!(projected[0].thinking.as_deref(), Some("complete reasoning"));
|
||||
let calls = projected[0].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(calls[0]["id"], "call-first");
|
||||
assert_eq!(calls[1]["id"], "call-second");
|
||||
assert_eq!(projected[1].tool_call_id.as_deref(), Some("call-first"));
|
||||
assert_eq!(projected[2].tool_call_id.as_deref(), Some("call-second"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(assets.mode(Mode::Agent).tools.len(), 20);
|
||||
assert_eq!(assets.mode(Mode::Ask).tools.len(), 18);
|
||||
assert_eq!(assets.mode(Mode::Plan).tools.len(), 16);
|
||||
assert_eq!(assets.mode(Mode::Debug).tools.len(), 18);
|
||||
assert_eq!(assets.mode(Mode::Multitask).tools.len(), 20);
|
||||
assert_eq!(assets.mode(Mode::Subagent).tools.len(), 4);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_tools_are_appended_after_the_stable_mode_tool_prefix() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let messages = vec![fixtures::user("u1", "one")];
|
||||
let base = compiler
|
||||
.compile(Mode::Agent, "model", "call", &messages)
|
||||
.unwrap();
|
||||
let dynamic = compiler
|
||||
.compile_with_dynamic_tools(
|
||||
Mode::Agent,
|
||||
"model",
|
||||
"call",
|
||||
&messages,
|
||||
&[ToolDefinition {
|
||||
name: "mcp_repo_lookup".into(),
|
||||
description: "lookup".into(),
|
||||
input_schema: serde_json::json!({"type": "object"}),
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(base.tools, dynamic.tools[..base.tools.len()]);
|
||||
assert_eq!(dynamic.tools.last().unwrap().name, "mcp_repo_lookup");
|
||||
}
|
||||
|
||||
fn tool_result(id: &str, output: serde_json::Value) -> CanonicalMessage {
|
||||
tool_result_with_call(id, &format!("call-{id}"), output)
|
||||
}
|
||||
|
||||
fn tool_result_with_call(
|
||||
id: &str,
|
||||
call_id: &str,
|
||||
output: impl Into<serde_json::Value>,
|
||||
) -> CanonicalMessage {
|
||||
CanonicalMessage {
|
||||
message_id: id.into(),
|
||||
role: Role::Tool,
|
||||
origin: Origin::Tool,
|
||||
content: MessageContent::ToolResult(ToolResultContent {
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
output: output.into(),
|
||||
is_error: false,
|
||||
}),
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_tool_pair(
|
||||
id: &str,
|
||||
model_call_id: &str,
|
||||
index: usize,
|
||||
call_id: &str,
|
||||
text: &str,
|
||||
thinking: &str,
|
||||
) -> CanonicalMessage {
|
||||
CanonicalMessage {
|
||||
message_id: id.into(),
|
||||
role: Role::Assistant,
|
||||
origin: Origin::Assistant,
|
||||
content: MessageContent::Assistant {
|
||||
text: text.into(),
|
||||
thinking: thinking.into(),
|
||||
model_call_id: Some(model_call_id.into()),
|
||||
tool_calls: vec![ToolCallContent {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
arguments: serde_json::json!({}),
|
||||
}],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use cursor_server::model::RuntimeEvent;
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_event_is_appended_exactly_once() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let event = RuntimeEvent {
|
||||
event_id: "branch:changed:7".into(),
|
||||
text: "runtime state changed".into(),
|
||||
};
|
||||
assert!(store
|
||||
.append_runtime_event_once("conversation", event.clone())
|
||||
.await
|
||||
.unwrap());
|
||||
assert!(!store
|
||||
.append_runtime_event_once("conversation", event)
|
||||
.await
|
||||
.unwrap());
|
||||
let messages = store.load_messages("conversation").await.unwrap();
|
||||
assert_eq!(messages.len(), 1);
|
||||
assert_eq!(
|
||||
messages[0].runtime_event_id.as_deref(),
|
||||
Some("branch:changed:7")
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
use bytes::Bytes;
|
||||
use cursor_server::{cursor::connect, Result};
|
||||
use prost::Message;
|
||||
|
||||
pub fn decode_single<M: Message + Default>(frame: &Bytes) -> Result<M> {
|
||||
let frames = connect::decode_frames(frame)?;
|
||||
assert_eq!(frames.len(), 1);
|
||||
Ok(M::decode(frames[0].1.clone())?)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
#![allow(dead_code)]
|
||||
|
||||
use std::{
|
||||
collections::VecDeque,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use cursor_server::{
|
||||
prompting::ModelRequest,
|
||||
provider::{Provider, ProviderStream, ResponseEvent},
|
||||
Error,
|
||||
};
|
||||
use futures_util::stream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
type FakeResponse = Vec<Result<ResponseEvent, Error>>;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct FakeProvider {
|
||||
responses: Arc<Mutex<VecDeque<FakeResponse>>>,
|
||||
requests: Arc<Mutex<Vec<ModelRequest>>>,
|
||||
}
|
||||
|
||||
impl FakeProvider {
|
||||
pub fn push(&self, events: Vec<ResponseEvent>) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(events.into_iter().map(Ok).collect());
|
||||
}
|
||||
pub fn push_error(&self, error: Error) {
|
||||
self.responses.lock().unwrap().push_back(vec![Err(error)]);
|
||||
}
|
||||
pub fn requests(&self) -> Vec<ModelRequest> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Provider for FakeProvider {
|
||||
fn stream(&self, request: ModelRequest, _cancellation: CancellationToken) -> ProviderStream {
|
||||
self.requests.lock().unwrap().push(request);
|
||||
let events = self
|
||||
.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop_front()
|
||||
.expect("fake response configured");
|
||||
Box::pin(stream::iter(events))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
#![allow(dead_code)]
|
||||
|
||||
use cursor_server::{
|
||||
model::{CanonicalMessage, Origin, Role},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
pub async fn temp_store() -> (tempfile::TempDir, Store) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let url = format!("sqlite://{}", directory.path().join("test.db").display());
|
||||
let store = Store::connect(&url).await.unwrap();
|
||||
(directory, store)
|
||||
}
|
||||
|
||||
pub fn user(id: &str, text: &str) -> CanonicalMessage {
|
||||
CanonicalMessage::text(id, Role::User, Origin::User, text)
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{connect, proto::agent::v1 as pb},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
provider::{FinishReason, ResponseEvent},
|
||||
run::{RunCommand, RunRegistry},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ResponseEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ResponseEvent::ThinkingStart,
|
||||
ResponseEvent::ThinkingDelta("reason".into()),
|
||||
ResponseEvent::ThinkingEnd,
|
||||
ResponseEvent::TextStart,
|
||||
ResponseEvent::TextDelta("hello".into()),
|
||||
ResponseEvent::TextEnd,
|
||||
ResponseEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run("request", "conversation", "hello")),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut text = String::new();
|
||||
let mut thinking = String::new();
|
||||
let mut thinking_duration_ms = None;
|
||||
let mut saw_turn_ended = false;
|
||||
let mut checkpoints = 0;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
match update.message {
|
||||
Some(pb::interaction_update::Message::TextDelta(delta)) => {
|
||||
text.push_str(&delta.text)
|
||||
}
|
||||
Some(pb::interaction_update::Message::ThinkingDelta(delta)) => {
|
||||
assert_eq!(
|
||||
delta.thinking_style,
|
||||
Some(pb::ThinkingStyle::Default as i32)
|
||||
);
|
||||
thinking.push_str(&delta.text);
|
||||
}
|
||||
Some(pb::interaction_update::Message::ThinkingCompleted(completed)) => {
|
||||
thinking_duration_ms = Some(completed.thinking_duration_ms)
|
||||
}
|
||||
Some(pb::interaction_update::Message::TurnEnded(_)) => saw_turn_ended = true,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(_)) => {
|
||||
checkpoints += 1
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert_eq!(text, "hello");
|
||||
assert_eq!(thinking, "reason");
|
||||
assert!(thinking_duration_ms.is_some_and(|duration| duration >= 1));
|
||||
assert!(saw_turn_ended);
|
||||
assert_eq!(checkpoints, 2, "final checkpoint is intentionally repeated");
|
||||
assert_eq!(provider.requests().len(), 1);
|
||||
assert!(store
|
||||
.load_messages("conversation")
|
||||
.await
|
||||
.unwrap()
|
||||
.iter()
|
||||
.any(|message| message.message_id == "user"));
|
||||
}
|
||||
|
||||
fn client_run(request_id: &str, conversation_id: &str, text: &str) -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: text.into(),
|
||||
message_id: "user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some(conversation_id.into()),
|
||||
run_id: Some(request_id.into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn kv_ack(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id,
|
||||
message: Some(pb::kv_client_message::Message::SetBlobResult(
|
||||
pb::SetBlobResult { error: None },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,521 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect, exec,
|
||||
pending::{ExecContext, PendingExecRegistry},
|
||||
proto::agent::v1 as pb,
|
||||
},
|
||||
model::{MessageContent, ToolCall},
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
provider::{FinishReason, ResponseEvent},
|
||||
run::{RunCommand, RunRegistry},
|
||||
};
|
||||
use prost::Message;
|
||||
use serde_json::json;
|
||||
|
||||
fn call(id: &str, name: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: id.into(),
|
||||
model_call_id: "model:0".into(),
|
||||
name: name.into(),
|
||||
arguments_text: "{}".into(),
|
||||
arguments: json!({}),
|
||||
}
|
||||
}
|
||||
|
||||
fn exec_context() -> ExecContext {
|
||||
ExecContext {
|
||||
conversation_id: "conversation".into(),
|
||||
terminals_folder: "/tmp/terminals".into(),
|
||||
admin_command_denylist: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "mcp-call".into(),
|
||||
model_call_id: "model:0".into(),
|
||||
name: "mcp_repo_lookup".into(),
|
||||
arguments_text: "{\"query\":\"x\"}".into(),
|
||||
arguments: json!({"query": "x"}),
|
||||
};
|
||||
let definition = pb::McpToolDefinition {
|
||||
name: "mcp_repo_lookup".into(),
|
||||
provider_identifier: "repo".into(),
|
||||
tool_name: "lookup".into(),
|
||||
description: "lookup".into(),
|
||||
input_schema: None,
|
||||
input_schema_json: None,
|
||||
};
|
||||
let message = cursor_server::cursor::exec::mcp_request(7, &call, &definition).unwrap();
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message else {
|
||||
panic!("expected ExecServerMessage")
|
||||
};
|
||||
let Some(pb::exec_server_message::Message::McpArgs(args)) = exec.message else {
|
||||
panic!("expected McpArgs")
|
||||
};
|
||||
assert_eq!(exec.exec_id, "mcp-call");
|
||||
assert_eq!(args.provider_identifier, "repo");
|
||||
assert_eq!(args.tool_name, "lookup");
|
||||
assert_eq!(
|
||||
args.args["query"].kind,
|
||||
Some(prost_types::value::Kind::StringValue("x".into()))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shell_uses_background_timeout_and_preserves_stream_identity() {
|
||||
let mut shell = call("call-shell", "Shell");
|
||||
shell.arguments = json!({
|
||||
"command": "python3 -m http.server 8000",
|
||||
"working_directory": "/tmp/project",
|
||||
"block_until_ms": 3000,
|
||||
"description": "Start HTTP server"
|
||||
});
|
||||
let context = exec_context();
|
||||
let request = exec::request(7, &shell, &context).unwrap();
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(request)) = request.message
|
||||
else {
|
||||
panic!("expected ExecServerMessage")
|
||||
};
|
||||
assert_eq!(request.accept_hook_additional_contexts, Some(true));
|
||||
let Some(pb::exec_server_message::Message::ShellStreamArgs(args)) = request.message else {
|
||||
panic!("expected ShellArgs")
|
||||
};
|
||||
assert_eq!(args.timeout, 3000);
|
||||
assert_eq!(
|
||||
args.timeout_behavior,
|
||||
pb::TimeoutBehavior::Background as i32
|
||||
);
|
||||
assert_eq!(args.hard_timeout, Some(86_400_000));
|
||||
assert_eq!(args.description.as_deref(), Some("Start HTTP server"));
|
||||
assert!(args.close_stdin);
|
||||
assert_eq!(args.conversation_id.as_deref(), Some("conversation"));
|
||||
assert_eq!(args.file_output_threshold_bytes, Some(40_000));
|
||||
|
||||
let pending = PendingExecRegistry::default();
|
||||
let id = pending.reserve(&shell, &context).await.unwrap();
|
||||
let delta = exec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::ShellStream(
|
||||
pb::ShellStream {
|
||||
event: Some(pb::shell_stream::Event::Stdout(pb::ShellStreamStdout {
|
||||
data: "Serving HTTP on port 8000\n".into(),
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
&pending,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let exec::ClientExecEvent::Delta(delta) = delta else {
|
||||
panic!("expected Shell stdout delta")
|
||||
};
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(delta)) = delta.message else {
|
||||
panic!("expected InteractionUpdate")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::ToolCallDelta(delta)) = delta.message else {
|
||||
panic!("expected ToolCallDelta")
|
||||
};
|
||||
assert_eq!(delta.call_id, "call-shell");
|
||||
assert_eq!(delta.model_call_id, "model:0");
|
||||
let Some(pb::tool_call_delta::Delta::ShellToolCallDelta(shell_delta)) =
|
||||
delta.tool_call_delta.and_then(|delta| delta.delta)
|
||||
else {
|
||||
panic!("expected ShellToolCallDelta")
|
||||
};
|
||||
let Some(pb::shell_tool_call_delta::Delta::Stdout(stdout)) = shell_delta.delta else {
|
||||
panic!("expected stdout")
|
||||
};
|
||||
assert_eq!(stdout.content, "Serving HTTP on port 8000\n");
|
||||
|
||||
let completion = exec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::ShellStream(
|
||||
pb::ShellStream {
|
||||
event: Some(pb::shell_stream::Event::Backgrounded(
|
||||
pb::ShellStreamBackgrounded {
|
||||
shell_id: 42,
|
||||
command: "python3 -m http.server 8000".into(),
|
||||
working_directory: "/tmp/project".into(),
|
||||
pid: Some(1234),
|
||||
ms_to_wait: Some(3000),
|
||||
reason: Some(pb::ShellBackgroundReason::Timeout as i32),
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
&pending,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let exec::ClientExecEvent::Completed(completion) = completion else {
|
||||
panic!("expected background completion")
|
||||
};
|
||||
assert_eq!(
|
||||
completion.result().output.as_str(),
|
||||
Some(
|
||||
"shell running in background shell_id=42 pid=1234 terminals_folder=/tmp/terminals\nServing HTTP on port 8000\n"
|
||||
)
|
||||
);
|
||||
let Some(pb::tool_call::Tool::ShellToolCall(tool)) = &completion.tool_call().tool else {
|
||||
panic!("expected ShellToolCall")
|
||||
};
|
||||
let result = tool.result.as_ref().expect("background ShellResult");
|
||||
assert_eq!(result.is_background, Some(true));
|
||||
assert_eq!(result.terminals_folder.as_deref(), Some("/tmp/terminals"));
|
||||
assert_eq!(result.pid, Some(1234));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn exec_ids_are_monotonic_and_released_ids_are_not_reused() {
|
||||
let pending = PendingExecRegistry::default();
|
||||
let first = pending
|
||||
.reserve(&call("call-1", "Read"), &exec_context())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(first, 1);
|
||||
assert_eq!(
|
||||
pending.call(first).await.map(|call| call.call_id),
|
||||
Some("call-1".into())
|
||||
);
|
||||
pending.discard(first).await;
|
||||
assert!(pending.call(first).await.is_none());
|
||||
|
||||
let second = pending
|
||||
.reserve(&call("call-2", "Read"), &exec_context())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(second, 2, "released Exec ids must not be reused in one Run");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_exec_client_message_is_not_a_terminal_result() {
|
||||
let pending = PendingExecRegistry::default();
|
||||
let id = pending
|
||||
.reserve(&call("call-1", "Read"), &exec_context())
|
||||
.await
|
||||
.unwrap();
|
||||
let event = exec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: None,
|
||||
..Default::default()
|
||||
},
|
||||
&pending,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(event, exec::ClientExecEvent::Pending));
|
||||
assert_eq!(
|
||||
pending.call(id).await.map(|call| call.call_id),
|
||||
Some("call-1".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tool_success_is_not_inferred_from_debug_text() {
|
||||
let pending = PendingExecRegistry::default();
|
||||
let mut write = call("call-1", "Write");
|
||||
write.arguments = json!({"path": "/tmp/a", "contents": "x"});
|
||||
let id = pending.reserve(&write, &exec_context()).await.unwrap();
|
||||
let event = exec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::WriteResult(
|
||||
pb::WriteResult {
|
||||
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
|
||||
path: "/tmp/a".into(),
|
||||
file_content_after_write: Some("enum Error { Example }".into()),
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
&pending,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let exec::ClientExecEvent::Completed(completion) = event else {
|
||||
panic!("expected terminal write result")
|
||||
};
|
||||
assert!(!completion.result().is_error);
|
||||
assert!(matches!(
|
||||
completion.tool_call().tool,
|
||||
Some(pb::tool_call::Tool::EditToolCall(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn an_exec_result_must_match_the_reserved_tool() {
|
||||
let pending = PendingExecRegistry::default();
|
||||
let id = pending
|
||||
.reserve(&call("call-1", "Read"), &exec_context())
|
||||
.await
|
||||
.unwrap();
|
||||
let result = exec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::WriteResult(
|
||||
pb::WriteResult {
|
||||
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
|
||||
path: "/tmp/a".into(),
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
&pending,
|
||||
)
|
||||
.await;
|
||||
let Err(error) = result else {
|
||||
panic!("mismatched result must fail")
|
||||
};
|
||||
assert!(error
|
||||
.to_string()
|
||||
.contains("unexpected Exec result for tool Read"));
|
||||
assert!(pending.call(id).await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
|
||||
let (directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ResponseEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ResponseEvent::ToolCallEnd { index: 0 },
|
||||
ResponseEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
provider.push(vec![
|
||||
ResponseEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ResponseEvent::TextStart,
|
||||
ResponseEvent::TextDelta("done".into()),
|
||||
ResponseEvent::TextEnd,
|
||||
ResponseEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("../prompt")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = RunRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
"test-model".into(),
|
||||
);
|
||||
let handle = registry.get_or_create("tool-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut seqno = 1;
|
||||
let mut saw_exec = false;
|
||||
let mut saw_typed_completion = false;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<serde_json::Value>(&payload).unwrap(),
|
||||
json!({})
|
||||
);
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
saw_exec = true;
|
||||
let exec_id = exec.id;
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: exec_id,
|
||||
exec_id: String::new(),
|
||||
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||
pb::ReadResult {
|
||||
result: Some(pb::read_result::Result::Success(
|
||||
pb::ReadSuccess {
|
||||
path: "/tmp/a".into(),
|
||||
total_lines: 1,
|
||||
file_size: 1,
|
||||
output: Some(
|
||||
pb::read_success::Output::Content(
|
||||
"x".into(),
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
handle
|
||||
.command(RunCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(
|
||||
pb::agent_client_message::Message::ExecClientControlMessage(
|
||||
pb::ExecClientControlMessage {
|
||||
message: Some(
|
||||
pb::exec_client_control_message::Message::StreamClose(
|
||||
pb::ExecClientStreamClose { id: exec_id },
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
if let Some(pb::interaction_update::Message::ToolCallCompleted(completed)) =
|
||||
update.message
|
||||
{
|
||||
let tool_call = completed.tool_call.expect("completed ToolCall");
|
||||
assert!(tool_call.started_at_ms.unwrap_or_default() > 1);
|
||||
assert!(tool_call.completed_at_ms.unwrap_or_default() > 1);
|
||||
assert!(tool_call.completed_at_ms >= tool_call.started_at_ms);
|
||||
let Some(pb::tool_call::Tool::ReadToolCall(read)) = tool_call.tool else {
|
||||
panic!("expected completed ReadToolCall")
|
||||
};
|
||||
let result = read.result.expect("typed ReadToolResult");
|
||||
assert!(matches!(
|
||||
result.result,
|
||||
Some(pb::read_tool_result::Result::Success(_))
|
||||
));
|
||||
saw_typed_completion = true;
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert!(saw_exec);
|
||||
assert!(saw_typed_completion);
|
||||
assert_eq!(provider.requests().len(), 2);
|
||||
let database = sqlx::SqlitePool::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let provider_call_index: i64 =
|
||||
sqlx::query_scalar("SELECT provider_call_index FROM runs WHERE request_id = ?")
|
||||
.bind("tool-request")
|
||||
.fetch_one(&database)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(provider_call_index, 1);
|
||||
let messages = store.load_messages("tool-conversation").await.unwrap();
|
||||
let result_position = messages
|
||||
.iter()
|
||||
.position(|message| matches!(message.content, MessageContent::ToolResult(_)))
|
||||
.expect("tool result persisted");
|
||||
let MessageContent::Assistant { tool_calls, .. } = &messages[result_position - 1].content
|
||||
else {
|
||||
panic!("tool result must immediately follow its assistant tool call")
|
||||
};
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].call_id, "call-1");
|
||||
}
|
||||
|
||||
fn client_run() -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: "read it".into(),
|
||||
message_id: "user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("tool-conversation".into()),
|
||||
run_id: Some("tool-request".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn kv_ack(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id,
|
||||
message: Some(pb::kv_client_message::Message::SetBlobResult(
|
||||
pb::SetBlobResult { error: None },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user