mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 19:31:28 +08:00
refactor: rebuild desktop app with Tauri
This commit is contained in:
@@ -0,0 +1,441 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
CursorCommand, CursorSessionHandle, CursorSessionRegistry,
|
||||
},
|
||||
model::{ContentPart, MessageContent, ProjectedContent, Role},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
const FOLLOW_UP: &str = "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]`.";
|
||||
const SHELL_FOLLOW_UP: &str = "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.";
|
||||
|
||||
#[tokio::test]
|
||||
async fn background_subagent_completion_starts_a_simulated_parent_turn() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(stop_response("model-call", "followed up"));
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("completion-request").await.unwrap();
|
||||
let (checkpoint, blobs) = drive_completion(
|
||||
&handle,
|
||||
completion_run(
|
||||
"child-id",
|
||||
"reusable-parent-run",
|
||||
pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Multitask as i32),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 1);
|
||||
let [runtime] = requests[0].history.as_slice() else {
|
||||
panic!("completion Run must add exactly one runtime message")
|
||||
};
|
||||
assert_eq!(runtime.role, Role::User);
|
||||
let ProjectedContent::Parts(parts) = &runtime.content else {
|
||||
panic!("completion context must be text")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("completion context must have one text part")
|
||||
};
|
||||
assert!(text.contains("kind: subagent"));
|
||||
assert!(text.contains("agent_id: child-id"));
|
||||
assert!(text.contains("child result"));
|
||||
assert!(text.contains(FOLLOW_UP));
|
||||
|
||||
let messages = store
|
||||
.load_current_messages(&cursor_server::model::ConversationId::new(
|
||||
"parent-conversation",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(messages.iter().any(|message| {
|
||||
message.runtime_event_id.as_deref() == Some("run-request:completion-request")
|
||||
&& matches!(&message.content, MessageContent::Parts { parts } if !parts.is_empty())
|
||||
}));
|
||||
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
blobs
|
||||
.get(checkpoint.turns.last().expect("completion Turn"))
|
||||
.expect("completion Turn Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent conversation Turn")
|
||||
};
|
||||
let user = pb::UserMessage::decode(
|
||||
blobs
|
||||
.get(&turn.user_message)
|
||||
.expect("simulated UserMessage Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(user.text.contains(FOLLOW_UP));
|
||||
assert_eq!(user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
assert_eq!(
|
||||
user.simulated_message_metadata.unwrap().task_id.as_deref(),
|
||||
Some("child-id")
|
||||
);
|
||||
|
||||
provider.push(stop_response("model-call-2", "followed up again"));
|
||||
let second = registry
|
||||
.get_or_create("completion-request-2")
|
||||
.await
|
||||
.unwrap();
|
||||
drive_completion(
|
||||
&second,
|
||||
completion_run("child-id-2", "reusable-parent-run-2", checkpoint),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let runtime_ids = requests[1]
|
||||
.history
|
||||
.iter()
|
||||
.map(|message| message.message_id.as_str())
|
||||
.filter(|id| id.starts_with("runtime:"))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
runtime_ids,
|
||||
[
|
||||
"runtime:run-request:completion-request",
|
||||
"runtime:run-request:completion-request-2"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn background_shell_completion_wakes_the_parent_with_the_captured_notification() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(stop_response(
|
||||
"shell-wakeup",
|
||||
"The background server was stopped.",
|
||||
));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("shell-completion-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let (checkpoint, blobs) = drive_completion(
|
||||
&handle,
|
||||
shell_completion_run(pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
let [runtime] = requests[0].history.as_slice() else {
|
||||
panic!("Shell completion Run must add exactly one runtime message")
|
||||
};
|
||||
let ProjectedContent::Parts(parts) = &runtime.content else {
|
||||
panic!("Shell completion context must be text")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("Shell completion context must have one text part")
|
||||
};
|
||||
assert!(text.contains("<system_notification>"));
|
||||
assert!(text.contains("kind: shell"));
|
||||
assert!(text.contains("status: aborted"));
|
||||
assert!(text.contains("task_id: 977679"));
|
||||
assert!(text.contains("detail: terminated_by_user"));
|
||||
assert!(text.contains("output_path: /tmp/977679.txt"));
|
||||
assert!(text.contains(SHELL_FOLLOW_UP));
|
||||
assert!(text.starts_with("<timestamp>"));
|
||||
assert!(!text.contains("You are still in **Agent Mode**"));
|
||||
assert!(text.find("<system_notification>").unwrap() < text.find("<user_query>").unwrap());
|
||||
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
blobs
|
||||
.get(checkpoint.turns.last().expect("Shell completion Turn"))
|
||||
.expect("Shell completion Turn Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent conversation Turn")
|
||||
};
|
||||
let user = pb::UserMessage::decode(
|
||||
blobs
|
||||
.get(&turn.user_message)
|
||||
.expect("simulated Shell UserMessage Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(user.text, *text);
|
||||
assert_eq!(user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
let metadata = 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"));
|
||||
}
|
||||
|
||||
async fn drive_completion(
|
||||
handle: &CursorSessionHandle,
|
||||
message: pb::AgentClientMessage,
|
||||
) -> (pb::ConversationStateStructure, HashMap<Vec<u8>, Vec<u8>>) {
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(message),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut blobs = HashMap::new();
|
||||
let mut final_checkpoint = None;
|
||||
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::ExecServerMessage(exec)) => {
|
||||
assert_eq!(exec.id, 0);
|
||||
assert!(matches!(
|
||||
exec.message,
|
||||
Some(pb::exec_server_message::Message::RequestContextArgs(_))
|
||||
));
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 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: 0 },
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: 0,
|
||||
message: Some(
|
||||
pb::exec_client_message::Message::RequestContextResult(
|
||||
pb::RequestContextResult {
|
||||
result: Some(
|
||||
pb::request_context_result::Result::Success(
|
||||
pb::RequestContextSuccess {
|
||||
request_context: Some(
|
||||
pb::RequestContext::default(),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = &kv.message {
|
||||
blobs.insert(set.blob_id.clone(), set.blob_data.clone());
|
||||
}
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
|
||||
if state.pending_tool_calls.is_empty() =>
|
||||
{
|
||||
final_checkpoint = Some(state);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
(
|
||||
final_checkpoint.expect("settled completion checkpoint"),
|
||||
blobs,
|
||||
)
|
||||
}
|
||||
|
||||
fn completion_run(
|
||||
child_id: &str,
|
||||
run_id: &str,
|
||||
conversation_state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(
|
||||
pb::conversation_action::Action::BackgroundTaskCompletionAction(
|
||||
pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![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()),
|
||||
output_path: Some("/tmp/child.jsonl".into()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
subagent_id: Some(child_id.into()),
|
||||
tool_call_id: Some("task-call".into()),
|
||||
..Default::default()
|
||||
}],
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("parent-conversation".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(conversation_state),
|
||||
run_id: Some(run_id.into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn shell_completion_run(
|
||||
conversation_state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(
|
||||
pb::conversation_action::Action::BackgroundTaskCompletionAction(
|
||||
pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![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()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
tool_call_id: Some("shell-call".into()),
|
||||
..Default::default()
|
||||
}],
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("parent-conversation".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(conversation_state),
|
||||
run_id: Some("shell-parent-run".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn stop_response(model_call_id: &str, text: &str) -> Vec<ModelEvent> {
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: model_call_id.into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta(text.into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]
|
||||
}
|
||||
|
||||
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,489 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashSet, sync::Arc};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
CursorCommand, CursorSessionRegistry,
|
||||
},
|
||||
model::ToolRoundId,
|
||||
provider::{FinishReason, ModelEvent},
|
||||
store::{BlobEdge, BlobId},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_dependencies_are_content_addressed_without_a_persistent_stream_outbox() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let child = store.put_blob(b"message", &[]).await.unwrap();
|
||||
let root = store
|
||||
.put_blob(
|
||||
b"checkpoint",
|
||||
&[BlobEdge {
|
||||
child: child.clone(),
|
||||
field_name: "turns[0]".into(),
|
||||
}],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(root, BlobId::digest(b"checkpoint"));
|
||||
assert_eq!(store.get_blob(&child).await.unwrap().unwrap(), b"message");
|
||||
let closure = store
|
||||
.blob_closure(std::slice::from_ref(&root))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(closure.contains(&root));
|
||||
assert!(closure.contains(&child));
|
||||
|
||||
let outbox: i64 = sqlx::query_scalar(
|
||||
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'outbox'",
|
||||
)
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(outbox, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn eligible_pending_checkpoint_resumes_tools_before_the_next_model_call() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-1".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "read-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-2".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("resumed".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
|
||||
let first = registry.get_or_create("first-run").await.unwrap();
|
||||
let mut first_output = first.subscribe();
|
||||
first
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(start_request()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut first_seqno = 1;
|
||||
let mut sent_blob_ids = HashSet::new();
|
||||
let staged = loop {
|
||||
let server = next_message(&mut first_output).await;
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
sent_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
acknowledge(&first, &mut first_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
|
||||
if state.pending_tool_calls.len() == 1 =>
|
||||
{
|
||||
break state;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
assert!(
|
||||
!sent_blob_ids.contains(
|
||||
BlobId::digest(&staged.encode_to_vec())
|
||||
.as_bytes()
|
||||
.as_slice()
|
||||
),
|
||||
"ConversationStateStructure is inline and must not be sent as a Blob"
|
||||
);
|
||||
let staged_started_at_ms =
|
||||
serde_json::from_str::<serde_json::Value>(staged.pending_tool_calls.first().unwrap())
|
||||
.unwrap()["providerOptions"]["cursor"]["pendingToolCallStartedAtMs"]
|
||||
.as_u64()
|
||||
.unwrap();
|
||||
assert_eq!(provider.requests().len(), 1);
|
||||
first.cancel();
|
||||
|
||||
let resumed = registry.get_or_create("resumed-run").await.unwrap();
|
||||
let mut resumed_output = resumed.subscribe();
|
||||
let mut resumed_checkpoints = Vec::new();
|
||||
let mut resumed_set_blob_ids = HashSet::new();
|
||||
resumed
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(resume_request(staged.clone())),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut resumed_seqno = 1;
|
||||
let exec_id = loop {
|
||||
let server = next_message(&mut resumed_output).await;
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
resumed_set_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
acknowledge(&resumed, &mut resumed_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
resumed_checkpoints.push(state);
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => break exec.id,
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
assert_eq!(
|
||||
provider.requests().len(),
|
||||
1,
|
||||
"resume must execute the pending batch before calling the model"
|
||||
);
|
||||
let resumed_round = store
|
||||
.tool_round(&ToolRoundId::new("resumed-run:round:resume"))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(resumed_round.created_at_ms, staged_started_at_ms);
|
||||
resumed
|
||||
.command(CursorCommand::Append {
|
||||
seqno: resumed_seqno,
|
||||
message: Box::new(read_result(exec_id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
resumed_seqno += 1;
|
||||
|
||||
let mut saw_settled_barrier_blob = false;
|
||||
let mut saw_settled_checkpoint = false;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), resumed_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)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
resumed_set_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
if !saw_settled_barrier_blob {
|
||||
assert_eq!(
|
||||
provider.requests().len(),
|
||||
1,
|
||||
"the next model call must wait for the settled Blob barrier"
|
||||
);
|
||||
saw_settled_barrier_blob = true;
|
||||
}
|
||||
acknowledge(&resumed, &mut resumed_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
if state.pending_tool_calls.is_empty() {
|
||||
saw_settled_checkpoint = true;
|
||||
}
|
||||
resumed_checkpoints.push(state);
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update))
|
||||
if matches!(
|
||||
update.message,
|
||||
Some(pb::interaction_update::Message::TextDelta(_))
|
||||
) =>
|
||||
{
|
||||
assert!(
|
||||
saw_settled_checkpoint,
|
||||
"the next model round must not become visible before settled checkpoint"
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert!(saw_settled_barrier_blob);
|
||||
assert!(saw_settled_checkpoint);
|
||||
assert!(staged
|
||||
.root_prompt_messages_json
|
||||
.iter()
|
||||
.all(|id| !resumed_set_blob_ids.contains(id)));
|
||||
assert_eq!(provider.requests().len(), 2);
|
||||
assert!(resumed_checkpoints
|
||||
.last()
|
||||
.unwrap()
|
||||
.read_paths
|
||||
.iter()
|
||||
.any(|path| path == "/tmp/a"));
|
||||
|
||||
let mut previous_steps = Vec::new();
|
||||
let mut saw_completed_read = false;
|
||||
for state in resumed_checkpoints {
|
||||
let Some(turn_id) = state.turns.last() else {
|
||||
continue;
|
||||
};
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
store
|
||||
.get_blob(&BlobId::from_bytes(turn_id).unwrap())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent turn");
|
||||
};
|
||||
assert!(turn.steps.len() >= previous_steps.len());
|
||||
assert_eq!(
|
||||
previous_steps,
|
||||
turn.steps[..previous_steps.len()],
|
||||
"published Step BlobIDs must be an immutable prefix"
|
||||
);
|
||||
previous_steps = turn.steps.clone();
|
||||
for step_id in &turn.steps {
|
||||
let step = pb::ConversationStep::decode(
|
||||
store
|
||||
.get_blob(&BlobId::from_bytes(step_id).unwrap())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
if let Some(pb::conversation_step::Message::ToolCall(call)) = step.message {
|
||||
if call.tool_call_id.as_deref() == Some("read-1") {
|
||||
assert!(call.started_at_ms.is_some());
|
||||
assert!(call.completed_at_ms.is_some());
|
||||
saw_completed_read = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
saw_completed_read,
|
||||
"settled Turn must keep the typed result"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn recovery_rejects_a_kv_get_payload_whose_hash_does_not_match_the_blob_id() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("bad-blob-run").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
let expected = BlobId::digest(b"expected");
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(resume_request(pb::ConversationStateStructure {
|
||||
root_prompt_messages_json: vec![expected.as_bytes().to_vec()],
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
})),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let get_id = loop {
|
||||
let server = next_message(&mut output).await;
|
||||
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = server.message {
|
||||
if matches!(
|
||||
kv.message,
|
||||
Some(pb::kv_server_message::Message::GetBlobArgs(_))
|
||||
) {
|
||||
break kv.id;
|
||||
}
|
||||
}
|
||||
};
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 1,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id: get_id,
|
||||
message: Some(pb::kv_client_message::Message::GetBlobResult(
|
||||
pb::GetBlobResult {
|
||||
blob_data: Some(b"corrupt".to_vec()),
|
||||
error: None,
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let error = 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 serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
|
||||
}
|
||||
};
|
||||
assert_eq!(error["error"]["code"], "invalid_argument");
|
||||
assert!(error["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("Blob hash mismatch"));
|
||||
}
|
||||
|
||||
fn start_request() -> 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: "user-1".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("conversation".into()),
|
||||
run_id: Some("first-run".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn resume_request(state: pb::ConversationStateStructure) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::ResumeAction(
|
||||
pb::ResumeAction::default(),
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(state),
|
||||
conversation_id: Some("conversation".into()),
|
||||
run_id: Some("resumed-run".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn read_result(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||
pb::ReadResult {
|
||||
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
|
||||
path: "/tmp/a".into(),
|
||||
output: Some(pb::read_success::Output::Content("value".into())),
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn next_message(
|
||||
output: &mut tokio::sync::mpsc::UnboundedReceiver<bytes::Bytes>,
|
||||
) -> pb::AgentServerMessage {
|
||||
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();
|
||||
assert_eq!(
|
||||
flags & connect::END_STREAM_FLAG,
|
||||
0,
|
||||
"unexpected EndStream: {}",
|
||||
String::from_utf8_lossy(&payload)
|
||||
);
|
||||
pb::AgentServerMessage::decode(payload).unwrap()
|
||||
}
|
||||
|
||||
async fn acknowledge(
|
||||
handle: &cursor_server::cursor::CursorSessionHandle,
|
||||
seqno: &mut i64,
|
||||
id: u32,
|
||||
) {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: *seqno,
|
||||
message: Box::new(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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
*seqno += 1;
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
client::{session, ClientCommand, ClientEvent, CommitCause},
|
||||
model::{
|
||||
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind,
|
||||
ToolDefinition, ToolResult,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
run::{RunEngine, RunOutcome},
|
||||
};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_client_without_checkpoint_protocol_runs_the_same_text_loop() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "call-1".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("hello".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let prepared = prepared(&store).await;
|
||||
let (port, mut client) = session(32);
|
||||
let engine = RunEngine::new(store.clone(), Arc::new(provider));
|
||||
let run =
|
||||
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
|
||||
|
||||
let mut saw_final_commit = false;
|
||||
while let Some(event) = client.events.recv().await {
|
||||
match event {
|
||||
ClientEvent::TextDelta(text) => assert_eq!(text, "hello"),
|
||||
ClientEvent::StateCommitted(state) => {
|
||||
saw_final_commit |= state.cause == CommitCause::FinalTurn;
|
||||
state.barrier.complete(Ok(()));
|
||||
}
|
||||
ClientEvent::Ended(outcome) => {
|
||||
assert_eq!(outcome, RunOutcome::Completed);
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert!(saw_final_commit);
|
||||
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_failed_claim_cannot_overwrite_the_existing_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let prepared = prepared(&store).await;
|
||||
store.claim_run(&prepared).await.unwrap();
|
||||
let (port, mut client) = session(8);
|
||||
let outcome = RunEngine::new(
|
||||
store.clone(),
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
)
|
||||
.run(prepared, port, CancellationToken::new())
|
||||
.await;
|
||||
|
||||
assert!(matches!(outcome, RunOutcome::Failed(_)));
|
||||
assert!(matches!(
|
||||
client.events.recv().await,
|
||||
Some(ClientEvent::Ended(RunOutcome::Failed(_)))
|
||||
));
|
||||
let status: String = sqlx::query_scalar("SELECT status FROM runs WHERE run_id = 'run'")
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(status, "running");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn required_client_state_failure_prevents_a_completed_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "call-1".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("hello".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let prepared = prepared(&store).await;
|
||||
let (port, mut client) = session(32);
|
||||
let engine = RunEngine::new(store, Arc::new(provider));
|
||||
let run =
|
||||
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
|
||||
|
||||
while let Some(event) = client.events.recv().await {
|
||||
match event {
|
||||
ClientEvent::StateCommitted(state) if state.cause == CommitCause::FinalTurn => {
|
||||
state.barrier.complete(Err("snapshot failed".into()));
|
||||
}
|
||||
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
|
||||
ClientEvent::Ended(outcome) => {
|
||||
assert_eq!(
|
||||
outcome,
|
||||
RunOutcome::Failed(cursor_server::run::RunFailure::Client(
|
||||
"snapshot failed".into()
|
||||
))
|
||||
);
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert_eq!(
|
||||
run.await.unwrap(),
|
||||
RunOutcome::Failed(cursor_server::run::RunFailure::Client(
|
||||
"snapshot failed".into()
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn generic_engine_waits_for_every_tool_result_without_any_cursor_wire_id() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "call-1".into(),
|
||||
},
|
||||
tool_start(0, "A"),
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
tool_start(1, "B"),
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 1,
|
||||
delta: "{}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 1 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "call-2".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("done".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let prepared = prepared(&store).await;
|
||||
let (port, mut client) = session(64);
|
||||
let commands = client.commands.clone();
|
||||
let engine = RunEngine::new(store, Arc::new(provider));
|
||||
let run =
|
||||
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
|
||||
while let Some(event) = client.events.recv().await {
|
||||
match event {
|
||||
ClientEvent::ExecuteToolRound { calls, .. } => {
|
||||
assert_eq!(calls.len(), 2);
|
||||
commands
|
||||
.send(ClientCommand::ToolResult(ToolResult {
|
||||
call_id: "B".into(),
|
||||
content: "result-B".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
commands
|
||||
.send(ClientCommand::ToolResult(ToolResult {
|
||||
call_id: "A".into(),
|
||||
content: "result-A".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
|
||||
ClientEvent::Ended(outcome) => {
|
||||
assert_eq!(outcome, RunOutcome::Completed);
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
|
||||
}
|
||||
|
||||
async fn prepared(store: &cursor_server::store::Store) -> PreparedRun {
|
||||
let conversation_id = ConversationId::new("conversation");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
PreparedRun {
|
||||
run_id: RunId::new("run"),
|
||||
conversation_id,
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new("model"),
|
||||
prompt: PromptSpec {
|
||||
instructions: "system".into(),
|
||||
tools: vec![ToolDefinition {
|
||||
name: "Tool".into(),
|
||||
description: "test".into(),
|
||||
parameters: serde_json::json!({"type":"object"}),
|
||||
}],
|
||||
},
|
||||
initial_messages: vec![fixtures::user("user", "hello")],
|
||||
action: RunAction::Start,
|
||||
base_revision_id: root,
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_start(index: usize, call_id: &str) -> ModelEvent {
|
||||
ModelEvent::ToolCallStart {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,372 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashMap, sync::Arc, time::Duration};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::{connect, proto::agent::v1 as pb, CursorCommand, CursorSessionRegistry},
|
||||
model::{
|
||||
ContentPart, ConversationId, MessageContent, Origin, ProjectedContent,
|
||||
ProviderEndpointInput, ProviderModelInput, ProviderType, Role, Usage,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn summarize_replaces_model_history_and_preserves_cursor_history() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let endpoint = store
|
||||
.create_provider(&ProviderEndpointInput {
|
||||
name: "Test".into(),
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
base_url: "https://example.com/v1".into(),
|
||||
api_key: None,
|
||||
custom_headers: serde_json::json!({}),
|
||||
extra_params: serde_json::json!({}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let model = store
|
||||
.save_provider_model(
|
||||
endpoint.provider_id,
|
||||
&ProviderModelInput {
|
||||
model_id: "test-model".into(),
|
||||
display_name: "Test Model".into(),
|
||||
endpoint_type: ProviderType::OpenAiChat,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: None,
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(text_response("old answer", 4_000, 12));
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "summary-call".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("Durable ".into()),
|
||||
ModelEvent::TextDelta("summary".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(4_012),
|
||||
output_tokens: Some(9),
|
||||
total_tokens: Some(4_021),
|
||||
..Default::default()
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
provider.push(text_response("new answer", 900, 5));
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
|
||||
let first = run(
|
||||
®istry,
|
||||
"first",
|
||||
user_request(
|
||||
"conversation",
|
||||
"user-1",
|
||||
"remember alpha",
|
||||
&model.model_hash,
|
||||
None,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let first_state = first.checkpoints.last().unwrap().clone();
|
||||
let old_turns = first_state.turns.clone();
|
||||
let old_roots = first_state.root_prompt_messages_json.clone();
|
||||
assert!(old_roots.len() >= 3);
|
||||
|
||||
let compacted = run(
|
||||
®istry,
|
||||
"compact",
|
||||
summary_request("conversation", &model.model_hash, first_state),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(compacted.summary_started, 1);
|
||||
assert_eq!(compacted.summary, "Durable summary");
|
||||
assert_eq!(compacted.summary_completed, 1);
|
||||
assert_eq!(compacted.turn_ended, 1);
|
||||
assert_eq!(compacted.token_delta, 0);
|
||||
assert_eq!(compacted.checkpoints.len(), 3);
|
||||
assert!(compacted
|
||||
.checkpoints
|
||||
.windows(2)
|
||||
.all(|pair| pair[0] == pair[1]));
|
||||
|
||||
let compacted_state = compacted.checkpoints.last().unwrap();
|
||||
assert_eq!(compacted_state.root_prompt_messages_json.len(), 2);
|
||||
assert!(compacted_state.turns.starts_with(&old_turns));
|
||||
assert_eq!(compacted_state.turns.len(), old_turns.len() + 1);
|
||||
assert_eq!(compacted_state.self_summary_count, 1);
|
||||
let summary_id = compacted_state.summary.as_ref().unwrap();
|
||||
let summary = pb::ConversationSummary::decode(compacted.blobs[summary_id].as_slice()).unwrap();
|
||||
assert_eq!(summary.summary, "Durable summary");
|
||||
let archive_id = compacted_state.summary_archive.as_ref().unwrap();
|
||||
let archive =
|
||||
pb::ConversationSummaryArchive::decode(compacted.blobs[archive_id].as_slice()).unwrap();
|
||||
assert_eq!(archive.summary, "Durable summary");
|
||||
assert_eq!(archive.window_tail, 0);
|
||||
assert_eq!(archive.summarized_messages, old_roots[1..]);
|
||||
assert_eq!(
|
||||
archive.summary_message,
|
||||
*compacted_state.root_prompt_messages_json.last().unwrap()
|
||||
);
|
||||
|
||||
let stored = store
|
||||
.load_current_messages(&ConversationId::new("conversation"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(stored.len(), 1);
|
||||
assert_eq!(stored[0].origin, Origin::Runtime);
|
||||
assert_eq!(stored[0].role, Role::User);
|
||||
assert!(matches!(
|
||||
&stored[0].content,
|
||||
MessageContent::Parts { parts }
|
||||
if matches!(parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text == "<conversation_summary>\nDurable summary\n</conversation_summary>")
|
||||
));
|
||||
|
||||
let after = run(
|
||||
®istry,
|
||||
"after",
|
||||
user_request(
|
||||
"conversation",
|
||||
"user-2",
|
||||
"what remains?",
|
||||
&model.model_hash,
|
||||
Some(compacted_state.clone()),
|
||||
),
|
||||
)
|
||||
.await;
|
||||
assert!(after
|
||||
.checkpoints
|
||||
.last()
|
||||
.unwrap()
|
||||
.root_prompt_messages_json
|
||||
.starts_with(&compacted_state.root_prompt_messages_json));
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 3);
|
||||
assert!(requests[1].prompt.tools.is_empty());
|
||||
assert!(requests[1]
|
||||
.prompt
|
||||
.instructions
|
||||
.contains("compacting conversation history"));
|
||||
assert_eq!(requests[1].history.len(), 2);
|
||||
assert_eq!(requests[2].history.len(), 2);
|
||||
let ProjectedContent::Parts(summary_parts) = &requests[2].history[0].content else {
|
||||
panic!("first post-compaction message must be the summary")
|
||||
};
|
||||
assert!(
|
||||
matches!(summary_parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text.contains("Durable summary"))
|
||||
);
|
||||
let ProjectedContent::Parts(new_user_parts) = &requests[2].history[1].content else {
|
||||
panic!("second post-compaction message must be the new runtime user")
|
||||
};
|
||||
assert!(
|
||||
matches!(new_user_parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text.contains("what remains?") && !text.contains("remember alpha"))
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct Output {
|
||||
checkpoints: Vec<pb::ConversationStateStructure>,
|
||||
blobs: HashMap<Vec<u8>, Vec<u8>>,
|
||||
summary: String,
|
||||
summary_started: usize,
|
||||
summary_completed: usize,
|
||||
turn_ended: usize,
|
||||
token_delta: usize,
|
||||
}
|
||||
|
||||
async fn run(
|
||||
registry: &CursorSessionRegistry,
|
||||
request_id: &str,
|
||||
request: pb::AgentClientMessage,
|
||||
) -> Output {
|
||||
let handle = registry.get_or_create(request_id).await.unwrap();
|
||||
let mut receiver = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(request),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut append_seqno = 1;
|
||||
let mut output = Output::default();
|
||||
loop {
|
||||
let frame = tokio::time::timeout(Duration::from_secs(5), receiver.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
return output;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = kv.message {
|
||||
output.blobs.insert(set.blob_id, set.blob_data);
|
||||
}
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
output.checkpoints.push(state)
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
match update.message {
|
||||
Some(pb::interaction_update::Message::SummaryStarted(_)) => {
|
||||
output.summary_started += 1
|
||||
}
|
||||
Some(pb::interaction_update::Message::Summary(delta)) => {
|
||||
output.summary.push_str(&delta.summary)
|
||||
}
|
||||
Some(pb::interaction_update::Message::SummaryCompleted(_)) => {
|
||||
output.summary_completed += 1
|
||||
}
|
||||
Some(pb::interaction_update::Message::TurnEnded(_)) => output.turn_ended += 1,
|
||||
Some(pb::interaction_update::Message::TokenDelta(_)) => output.token_delta += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn text_response(text: &str, input: u64, output: u64) -> Vec<ModelEvent> {
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: format!("call-{text}"),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta(text.into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(input),
|
||||
output_tokens: Some(output),
|
||||
total_tokens: Some(input + output),
|
||||
..Default::default()
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]
|
||||
}
|
||||
|
||||
fn user_request(
|
||||
conversation_id: &str,
|
||||
message_id: &str,
|
||||
text: &str,
|
||||
model_id: &str,
|
||||
state: Option<pb::ConversationStateStructure>,
|
||||
) -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: text.into(),
|
||||
message_id: message_id.into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
request(
|
||||
conversation_id,
|
||||
model_id,
|
||||
state,
|
||||
pb::conversation_action::Action::UserMessageAction(pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
request_context: Some(pb::RequestContext::default()),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn summary_request(
|
||||
conversation_id: &str,
|
||||
model_id: &str,
|
||||
state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: "/summarize".into(),
|
||||
message_id: "summary-command".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
request(
|
||||
conversation_id,
|
||||
model_id,
|
||||
Some(state),
|
||||
pb::conversation_action::Action::UserMessageAction(pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
request_context: Some(pb::RequestContext::default()),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn request(
|
||||
conversation_id: &str,
|
||||
model_id: &str,
|
||||
state: Option<pb::ConversationStateStructure>,
|
||||
action: pb::conversation_action::Action,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: model_id.into(),
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(action),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some(conversation_id.into()),
|
||||
conversation_state: state,
|
||||
run_id: Some("reusable-wire-run-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,137 @@
|
||||
#[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::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::CursorSessionRegistry,
|
||||
cursor::{
|
||||
connect, handlers,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
};
|
||||
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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
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)
|
||||
.unwrap()
|
||||
.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,416 @@
|
||||
#[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::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::{
|
||||
connect,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
cursor::{CursorCommand, CursorSessionRegistry},
|
||||
model::{MessageContent, Role},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
Error,
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_error() {
|
||||
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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("failed-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut checkpoints = Vec::new();
|
||||
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(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
checkpoints.push(state);
|
||||
}
|
||||
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_eq!(
|
||||
checkpoints.len(),
|
||||
1,
|
||||
"the initial user state is checkpointed"
|
||||
);
|
||||
assert!(checkpoints[0].pending_tool_calls.is_empty());
|
||||
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
|
||||
);
|
||||
assert_eq!(decoded.is_expected, Some(true));
|
||||
let custom = decoded.details.unwrap();
|
||||
assert_eq!(custom.title, "Provider 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_current_messages(&cursor_server::model::ConversationId::new(
|
||||
"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![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("protocol-failed-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::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(CursorCommand::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(CursorCommand::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"],
|
||||
"unknown ExecClientMessage id: 1001"
|
||||
);
|
||||
assert_eq!(
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn duplicate_run_request_on_one_bidi_stream_is_a_protocol_error() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("protocol-failed-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(protocol_client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let mut duplicate_sent = 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(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(_)) if !duplicate_sent => {
|
||||
duplicate_sent = true;
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(protocol_client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
|
||||
assert!(duplicate_sent);
|
||||
assert_eq!(error_json["error"]["code"], "invalid_argument");
|
||||
assert_eq!(
|
||||
error_json["error"]["message"],
|
||||
"duplicate RunRequest for request_id: protocol-failed-request"
|
||||
);
|
||||
}
|
||||
|
||||
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()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..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()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..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,489 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::{connect, proto::agent::v1 as pb},
|
||||
cursor::{CursorCommand, CursorSessionRegistry},
|
||||
model::{ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
run::RunRegistry,
|
||||
store::RunStatus,
|
||||
};
|
||||
use prost::Message;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
#[tokio::test]
|
||||
async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
|
||||
let registry = RunRegistry::default();
|
||||
let conversation = cursor_server::model::ConversationId::new("conversation");
|
||||
let first = CancellationToken::new();
|
||||
let second = CancellationToken::new();
|
||||
registry
|
||||
.activate(
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("first"),
|
||||
first.clone(),
|
||||
)
|
||||
.await;
|
||||
registry
|
||||
.activate(
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("second"),
|
||||
second.clone(),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(first.is_cancelled());
|
||||
assert!(!second.is_cancelled());
|
||||
registry
|
||||
.release(&conversation, &cursor_server::model::RunId::new("first"))
|
||||
.await;
|
||||
registry.shutdown().await;
|
||||
assert!(second.is_cancelled());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_replaced_run_cannot_overwrite_its_cancelled_status() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("conversation");
|
||||
let base_revision_id = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let prepared = |run_id: &str| PreparedRun {
|
||||
run_id: RunId::new(run_id),
|
||||
conversation_id: conversation_id.clone(),
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new("model"),
|
||||
prompt: PromptSpec {
|
||||
instructions: String::new(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
initial_messages: Vec::new(),
|
||||
action: RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
},
|
||||
base_revision_id,
|
||||
};
|
||||
let first = prepared("first");
|
||||
let second = prepared("second");
|
||||
|
||||
store.claim_run(&first).await.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO llm_calls(
|
||||
call_id, run_id, conversation_id, provider_call_index,
|
||||
provider_type, provider_url, request_type, request_url,
|
||||
model_id, display_name, status,
|
||||
created_at_ms, message_count, tool_count, detailed
|
||||
) VALUES (
|
||||
'first:0', 'first', 'conversation', 0,
|
||||
'openai-chat', 'https://example.com/v1',
|
||||
'openai-chat', 'https://example.com/v1/chat/completions',
|
||||
'model', 'Model', 'running',
|
||||
unixepoch('subsec') * 1000, 1, 0, 0
|
||||
)",
|
||||
)
|
||||
.execute(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
store.claim_run(&second).await.unwrap();
|
||||
assert!(!store
|
||||
.finish_run(&first.run_id, RunStatus::Completed, None, None,)
|
||||
.await
|
||||
.unwrap());
|
||||
|
||||
let status: String = sqlx::query_scalar("SELECT status FROM runs WHERE run_id = 'first'")
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let active: Option<String> = sqlx::query_scalar(
|
||||
"SELECT active_run_id FROM conversations WHERE conversation_id = 'conversation'",
|
||||
)
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(status, "cancelled");
|
||||
assert_eq!(active.as_deref(), Some("second"));
|
||||
let call: (String, Option<i64>, Option<i64>) = sqlx::query_as(
|
||||
"SELECT status, finished_at_ms, duration_ms FROM llm_calls WHERE call_id = 'first:0'",
|
||||
)
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(call.0, "cancelled");
|
||||
assert!(call.1.is_some());
|
||||
assert!(call.2.is_some());
|
||||
}
|
||||
|
||||
#[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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
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 runtime_user_message_action_aborts_active_exec_before_canceled_end_stream() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{\"path\":\"/tmp/a\"}".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("cancel-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::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 (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
assert_eq!(
|
||||
flags & connect::END_STREAM_FLAG,
|
||||
0,
|
||||
"Run ended before Exec: {}",
|
||||
String::from_utf8_lossy(&payload)
|
||||
);
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(CursorCommand::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
|
||||
.command(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(runtime_user_message()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
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);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn injected_user_context_restarts_only_the_active_model_cycle() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push_pending();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "continued".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("continued after injection".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("inject-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run_for("inject-request", "inject-conversation")),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
|
||||
while provider.requests().is_empty() {
|
||||
assert!(
|
||||
tokio::time::Instant::now() < deadline,
|
||||
"provider did not start"
|
||||
);
|
||||
if let Ok(Some(frame)) =
|
||||
tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await
|
||||
{
|
||||
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
|
||||
}
|
||||
}
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(runtime_injection()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
|
||||
let mut protocol_events = Vec::new();
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before successful EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
assert_eq!(payload.as_ref(), b"{}");
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message {
|
||||
match update.message {
|
||||
Some(pb::interaction_update::Message::ContextInjectionState(update)) => {
|
||||
assert_eq!(update.injection_id, "injection-1");
|
||||
match update.state.and_then(|state| state.state) {
|
||||
Some(pb::context_injection_state::State::Queued(_)) => {
|
||||
protocol_events.push("queued")
|
||||
}
|
||||
Some(pb::context_injection_state::State::Delivered(delivered)) => {
|
||||
assert!(!delivered.delivery_batch_id.is_empty());
|
||||
assert!(delivered.delivered_at_ms > 0);
|
||||
protocol_events.push("delivered");
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Some(pb::interaction_update::Message::UserMessageAppended(update)) => {
|
||||
let user = update.user_message.expect("appended user message");
|
||||
assert_eq!(user.message_id, "injected-user");
|
||||
assert_eq!(user.text, "injected follow-up");
|
||||
protocol_events.push("user_message_appended");
|
||||
}
|
||||
Some(pb::interaction_update::Message::TextDelta(update))
|
||||
if update.text.contains("continued after injection") =>
|
||||
{
|
||||
protocol_events.push("continued_output");
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
|
||||
}
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let continued_history = serde_json::to_string(&requests[1].history).unwrap();
|
||||
assert!(continued_history.contains("injected follow-up"));
|
||||
assert!(!handle.cancellation().is_cancelled());
|
||||
assert_eq!(
|
||||
protocol_events,
|
||||
[
|
||||
"queued",
|
||||
"delivered",
|
||||
"user_message_appended",
|
||||
"continued_output"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
fn client_run() -> pb::AgentClientMessage {
|
||||
client_run_for("cancel-request", "cancel-conversation")
|
||||
}
|
||||
|
||||
fn client_run_for(request_id: &str, conversation_id: &str) -> 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(conversation_id.into()),
|
||||
run_id: Some(request_id.into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn acknowledge_kv(
|
||||
handle: &cursor_server::cursor::CursorSessionHandle,
|
||||
append_seqno: &mut i64,
|
||||
frame: &[u8],
|
||||
) {
|
||||
let (flags, payload) = connect::decode_frames(frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
return;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = server.message {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: *append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
*append_seqno += 1;
|
||||
}
|
||||
}
|
||||
|
||||
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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_user_message() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ConversationAction(
|
||||
pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "queued follow-up".into(),
|
||||
message_id: "queued-user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_injection() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ConversationAction(
|
||||
pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::InjectContextAction(
|
||||
pb::InjectContextAction {
|
||||
injection_id: "injection-1".into(),
|
||||
expected_run_id: "inject-request".into(),
|
||||
payload: Some(pb::inject_context_action::Payload::UserContext(
|
||||
pb::UserContextInjection {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "injected follow-up".into(),
|
||||
message_id: "injected-user".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
request_context: Some(Default::default()),
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::{http::header, response::IntoResponse, routing::post, Router};
|
||||
use cursor_server::{
|
||||
model::{
|
||||
ModelInvocation, ModelRequest, ModelSpec, PromptSpec, ProviderEndpointInput,
|
||||
ProviderModelInput, ProviderType,
|
||||
},
|
||||
provider::{ModelEvent, Provider, ProviderRouter},
|
||||
store::Store,
|
||||
};
|
||||
use futures_util::StreamExt;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
async fn test_store(name: &str) -> (tempfile::TempDir, Store) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join(name).display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
(directory, store)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cursor_traces_are_absent_when_detailed_logging_is_disabled() {
|
||||
let (_directory, store) = test_store("cursor-trace-disabled.db").await;
|
||||
assert!(!store
|
||||
.start_cursor_trace_if_detailed(
|
||||
"request-disabled",
|
||||
Some("conversation"),
|
||||
"local_byok",
|
||||
Some("model"),
|
||||
)
|
||||
.await
|
||||
.unwrap());
|
||||
assert!(store
|
||||
.cursor_trace("request-disabled")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cursor_trace_links_detailed_artifacts_to_the_logical_run() {
|
||||
let (_directory, store) = test_store("cursor-trace-enabled.db").await;
|
||||
store.set_detailed_logging(true).await.unwrap();
|
||||
assert!(store
|
||||
.start_cursor_trace_if_detailed(
|
||||
"request-enabled",
|
||||
Some("conversation"),
|
||||
"cursor_official",
|
||||
Some("official-model"),
|
||||
)
|
||||
.await
|
||||
.unwrap());
|
||||
store
|
||||
.append_cursor_trace_artifact(
|
||||
"request-enabled",
|
||||
"bidi_append_request",
|
||||
"cursor_client",
|
||||
b"request",
|
||||
&serde_json::json!({"append_seqno": 1}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.add_cursor_trace_request_bytes("request-enabled", 7)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.start_cursor_trace_response("request-enabled", 200)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.add_cursor_trace_response_chunk("request-enabled", "cursor_official", b"response")
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.finish_cursor_trace("request-enabled", None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let trace = store
|
||||
.cursor_trace("request-enabled")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(trace.route, "cursor_official");
|
||||
assert_eq!(trace.status, "completed");
|
||||
assert_eq!(trace.request_bytes, 7);
|
||||
assert_eq!(trace.response_bytes, 8);
|
||||
assert_eq!(trace.response_event_count, 1);
|
||||
let artifacts = store
|
||||
.cursor_trace_artifacts("request-enabled")
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(artifacts.len(), 2);
|
||||
assert_eq!(artifacts[0].artifact_type, "bidi_append_request");
|
||||
assert_eq!(artifacts[1].artifact_type, "run_sse_chunk");
|
||||
assert_eq!(artifacts[1].data, b"response");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn records_one_summary_and_raw_payloads_for_one_provider_request() {
|
||||
let app = Router::new().route(
|
||||
"/v1/chat/completions",
|
||||
post(|| async {
|
||||
(
|
||||
[(header::CONTENT_TYPE, "text/event-stream")],
|
||||
concat!(
|
||||
"data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":null}]}\n\n",
|
||||
"data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":2,\"total_tokens\":12}}\n\n",
|
||||
"data: [DONE]\n\n"
|
||||
),
|
||||
)
|
||||
.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, app).await.unwrap() });
|
||||
|
||||
let (_directory, store) = test_store("observability.db").await;
|
||||
store.set_detailed_logging(true).await.unwrap();
|
||||
let endpoint = store
|
||||
.create_provider(&ProviderEndpointInput {
|
||||
name: "test".into(),
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
base_url: format!("http://{address}/v1"),
|
||||
api_key: Some("not-recorded".into()),
|
||||
custom_headers: serde_json::json!({"x-safe":"visible","authorization":"hidden"}),
|
||||
extra_params: serde_json::json!({}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let model = store
|
||||
.save_provider_model(
|
||||
endpoint.provider_id,
|
||||
&ProviderModelInput {
|
||||
model_id: "actual-model".into(),
|
||||
display_name: "Display Model".into(),
|
||||
endpoint_type: ProviderType::OpenAiChat,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: None,
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let provider = ProviderRouter::new(store.clone(), Duration::from_secs(5));
|
||||
let events = provider
|
||||
.stream(
|
||||
ModelInvocation {
|
||||
call_id: "call-1".into(),
|
||||
run_id: "run-1".into(),
|
||||
conversation_id: "conversation-1".into(),
|
||||
provider_call_index: 0,
|
||||
request: ModelRequest {
|
||||
prompt: PromptSpec {
|
||||
instructions: "system".into(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
model: ModelSpec::new(model.model_hash),
|
||||
history: Vec::new(),
|
||||
},
|
||||
},
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
assert!(events.iter().all(Result::is_ok));
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, Ok(ModelEvent::Done(_)))));
|
||||
|
||||
let call = store.llm_call("call-1").await.unwrap().unwrap();
|
||||
assert_eq!(call.status, "completed");
|
||||
assert_eq!(call.request_type, "openai-chat");
|
||||
assert!(call.request_url.ends_with("/v1/chat/completions"));
|
||||
assert_eq!(call.total_tokens, Some(12));
|
||||
assert!(call.ttfb_ms.is_some());
|
||||
assert!(call.ttft_ms.is_some());
|
||||
let request = store.llm_call_request("call-1").await.unwrap().unwrap();
|
||||
assert_eq!(request.body["model"], "actual-model");
|
||||
assert_eq!(request.headers["x-safe"], "visible");
|
||||
assert!(request.headers.get("authorization").is_none());
|
||||
assert!(!store.llm_call_chunks("call-1").await.unwrap().is_empty());
|
||||
server.abort();
|
||||
}
|
||||
@@ -0,0 +1,583 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{Mode, PromptAssets, PromptCompiler},
|
||||
model::{project_messages, ProjectedContent},
|
||||
model::{
|
||||
CanonicalMessage, MessageContent, ModelSpec, Origin, Role, ToolCallContent, ToolDefinition,
|
||||
ToolResultContent,
|
||||
},
|
||||
};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
#[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 ProjectedContent::ToolResult(object_result) = &projected[0].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
let object_text = &object_result.content;
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(object_text).unwrap(),
|
||||
object
|
||||
);
|
||||
let ProjectedContent::ToolResult(string_result) = &projected[1].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
assert_eq!(string_result.content, "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(),
|
||||
tool_round_id: Some("round".into()),
|
||||
replay_state: None,
|
||||
tool_calls: Vec::new(),
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}];
|
||||
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
let ProjectedContent::Assistant { text, thinking, .. } = &projected[0].content else {
|
||||
panic!("expected assistant")
|
||||
};
|
||||
assert_eq!(text, "visible answer");
|
||||
assert_eq!(thinking, "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, Role::Assistant);
|
||||
let ProjectedContent::Assistant {
|
||||
thinking, calls, ..
|
||||
} = &projected[0].content
|
||||
else {
|
||||
panic!("expected assistant")
|
||||
};
|
||||
assert_eq!(thinking, "complete reasoning");
|
||||
assert_eq!(calls[0].call_id, "call-first");
|
||||
assert_eq!(calls[1].call_id, "call-second");
|
||||
let ProjectedContent::ToolResult(second) = &projected[1].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
let ProjectedContent::ToolResult(first) = &projected[2].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
assert_eq!(second.call_id, "call-second");
|
||||
assert_eq!(first.call_id, "call-first");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(assets.mode(Mode::Agent).tools.len(), 22);
|
||||
assert_eq!(
|
||||
assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"GenerateImage",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Glob",
|
||||
"AskQuestion",
|
||||
"Task",
|
||||
"AwaitShell",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
]
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Ask,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Plan,
|
||||
&[
|
||||
"Shell",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"TodoWrite",
|
||||
"ReadLints",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"AskQuestion",
|
||||
"CreatePlan",
|
||||
"Task",
|
||||
"FetchMcpResource",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"235a2a9a7785844eb5186f1c8f2294a36a04bbf103887d05ec386f8c7cc52abc",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Debug,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Multitask,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"SwitchMode",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"GenerateImage",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"04c5fb238eb3695936ceed610b481caf7507f934efb88cb7130a8756e12959e3",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Subagent,
|
||||
&[
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"GenerateImage",
|
||||
"ReadLints",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"Glob",
|
||||
"AwaitShell",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"UpdateCurrentStep",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"f88c55fdbb53be377e64cc6280ebb23c75e2d244b5c3463ed18e90752c3b7ff5",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Compaction,
|
||||
&[],
|
||||
"4f53cda18c2baa0c0354bb5f9a3ecbe5ed12ab4d8e11ba873c2f11161202b945",
|
||||
);
|
||||
assert_eq!(
|
||||
schema_digest(&assets.mode(Mode::Agent).tools),
|
||||
"4324c36fa047fbe4c93a5d5f0b736c559a942e097266c5bae057804289f8b359"
|
||||
);
|
||||
let shell = assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.find(|tool| tool.name == "Shell")
|
||||
.unwrap();
|
||||
assert!(
|
||||
shell.parameters["properties"]["block_until_ms"]["description"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("do not combine it with `nohup`, `&`, `disown`")
|
||||
);
|
||||
for mode in [
|
||||
Mode::Agent,
|
||||
Mode::Ask,
|
||||
Mode::Debug,
|
||||
Mode::Multitask,
|
||||
Mode::Subagent,
|
||||
Mode::Compaction,
|
||||
] {
|
||||
assert!(!assets
|
||||
.mode(mode)
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "CreatePlan" || tool.name == "PatchEdit"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_captured_mode_owns_and_renders_its_runtime_template() {
|
||||
let compiler = PromptCompiler::new(
|
||||
PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let values = BTreeMap::from([
|
||||
("REQUEST_CONTEXT", String::new()),
|
||||
("OPEN_FILES", String::new()),
|
||||
("SELECTED_CONTEXT", String::new()),
|
||||
("ACTION_CONTEXT", String::new()),
|
||||
("TIMESTAMP", "Sunday, Aug 16, 2026, 11:31 PM (UTC+8)".into()),
|
||||
("USER_QUERY", "question".into()),
|
||||
("DEBUG_SERVER_ENDPOINT", "http://debug".into()),
|
||||
("DEBUG_LOG_PATH", "/tmp/debug.log".into()),
|
||||
("DEBUG_SESSION_ID", "session".into()),
|
||||
]);
|
||||
for (mode, marker) in [
|
||||
(Mode::Agent, "You are still in **Agent Mode**"),
|
||||
(Mode::Ask, "Ask mode is active."),
|
||||
(Mode::Plan, "Plan mode is active."),
|
||||
(Mode::Debug, "You are now in **DEBUG MODE**"),
|
||||
(Mode::Multitask, "The user has engaged **Multitask Mode**"),
|
||||
] {
|
||||
let rendered = compiler.runtime_message(mode, &values).unwrap();
|
||||
assert!(rendered.contains(marker), "missing {mode:?} marker");
|
||||
assert!(rendered.contains("<user_query>\nquestion\n</user_query>"));
|
||||
assert_eq!(rendered.matches("<user_query>").count(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_mode(assets: &PromptAssets, mode: Mode, expected: &[&str], digest: &str) {
|
||||
assert_eq!(
|
||||
assets
|
||||
.mode(mode)
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
expected
|
||||
);
|
||||
assert_eq!(schema_digest(&assets.mode(mode).tools), digest);
|
||||
}
|
||||
|
||||
fn schema_digest(tools: &[ToolDefinition]) -> String {
|
||||
hex::encode(Sha256::digest(serde_json::to_vec(tools).unwrap()))
|
||||
}
|
||||
|
||||
#[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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let base = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let dynamic = compiler
|
||||
.prompt_spec(
|
||||
Mode::Agent,
|
||||
&ModelSpec::new("model"),
|
||||
&[ToolDefinition {
|
||||
name: "mcp_repo_lookup".into(),
|
||||
description: "lookup".into(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
}],
|
||||
false,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(base.tools, dynamic.tools[..base.tools.len()]);
|
||||
assert_eq!(dynamic.tools.last().unwrap().name, "mcp_repo_lookup");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_tool_cannot_replace_a_mode_tool() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let error = compiler
|
||||
.prompt_spec(
|
||||
Mode::Agent,
|
||||
&ModelSpec::new("model"),
|
||||
&[ToolDefinition {
|
||||
name: "Read".into(),
|
||||
description: "replacement".into(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
}],
|
||||
false,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(error
|
||||
.to_string()
|
||||
.contains("dynamic MCP tool conflicts with a mode tool: Read"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn image_generation_capability_controls_only_the_generate_image_definition() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let without = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let mut model = ModelSpec::new("model");
|
||||
model.supports_image_generation = true;
|
||||
let with = compiler
|
||||
.prompt_spec(Mode::Agent, &model, &[], false)
|
||||
.unwrap();
|
||||
|
||||
assert!(!without
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "GenerateImage"));
|
||||
assert!(with.tools.iter().any(|tool| tool.name == "GenerateImage"));
|
||||
assert_eq!(with.tools.len(), without.tools.len() + 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_system_prompt_is_static_and_substitutes_the_model_name() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let mut model = ModelSpec::new("test-model-hash");
|
||||
model.display_name = Some("Test Model".into());
|
||||
let request = compiler
|
||||
.prompt_spec(Mode::Agent, &model, &[], false)
|
||||
.unwrap();
|
||||
let prompt = &request.instructions;
|
||||
assert!(prompt.contains("powered by Test Model"));
|
||||
assert!(!prompt.contains("test-model-hash"));
|
||||
assert!(!prompt.contains("{{FAKE_MODEL_NAME}}"));
|
||||
assert!(!prompt.contains("<user_info>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subagent_uses_the_agent_prompt_and_only_the_captured_tool_delta() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let agent_prompt = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let subagent_prompt = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
assert_eq!(agent_prompt.instructions, subagent_prompt.instructions);
|
||||
|
||||
let request = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
request
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"ReadLints",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"Glob",
|
||||
"AwaitShell",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"UpdateCurrentStep",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
]
|
||||
);
|
||||
assert!(!request.tools.iter().any(|tool| tool.name == "Task"));
|
||||
|
||||
let suppressed = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], true)
|
||||
.unwrap();
|
||||
assert!(!suppressed
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "UpdateCurrentStep"));
|
||||
}
|
||||
|
||||
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 {
|
||||
let output = output.into();
|
||||
CanonicalMessage {
|
||||
message_id: id.into(),
|
||||
role: Role::Tool,
|
||||
origin: Origin::Tool,
|
||||
content: MessageContent::ToolResult(ToolResultContent {
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
content: output
|
||||
.as_str()
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(|| output.to_string()),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_tool_pair(
|
||||
id: &str,
|
||||
tool_round_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(),
|
||||
tool_round_id: Some(tool_round_id.into()),
|
||||
replay_state: None,
|
||||
tool_calls: vec![ToolCallContent {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
arguments: serde_json::json!({}),
|
||||
}],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
use cursor_server::{
|
||||
model::{
|
||||
ConversationId, ModelSpec, NewLlmCall, PreparedRun, PromptSpec, ProviderEndpointInput,
|
||||
ProviderModelInput, ProviderType, RunAction, RunId, RunKind, Usage,
|
||||
},
|
||||
store::{RunStatus, Store},
|
||||
};
|
||||
|
||||
async fn 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)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_secret_is_write_only_and_model_hash_is_stable() {
|
||||
let (_directory, store) = store().await;
|
||||
let provider = store
|
||||
.create_provider(&ProviderEndpointInput {
|
||||
name: "Local".into(),
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
base_url: "https://example.com/v1/".into(),
|
||||
api_key: Some("secret".into()),
|
||||
custom_headers: serde_json::json!({"x-route":"one", "authorization":"header-secret"}),
|
||||
extra_params: serde_json::json!({"temperature":0}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(provider.has_api_key);
|
||||
assert!(!serde_json::to_string(&provider).unwrap().contains("secret"));
|
||||
assert_eq!(
|
||||
provider.custom_headers["authorization"],
|
||||
serde_json::Value::Null
|
||||
);
|
||||
let updated = store
|
||||
.update_provider(
|
||||
provider.provider_id,
|
||||
&ProviderEndpointInput {
|
||||
name: "Renamed".into(),
|
||||
provider_type: provider.provider_type,
|
||||
base_url: provider.base_url.clone(),
|
||||
api_key: None,
|
||||
custom_headers: provider.custom_headers.clone(),
|
||||
extra_params: provider.extra_params.clone(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(updated.name, "Renamed");
|
||||
assert_eq!(
|
||||
store
|
||||
.provider(provider.provider_id)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.custom_headers["authorization"],
|
||||
"header-secret"
|
||||
);
|
||||
|
||||
let model = store
|
||||
.save_provider_model(
|
||||
provider.provider_id,
|
||||
&ProviderModelInput {
|
||||
model_id: "model-a".into(),
|
||||
display_name: "Model A".into(),
|
||||
endpoint_type: ProviderType::OpenAiChat,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: None,
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: true,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(model.model_hash, "f246010a");
|
||||
assert!(model.supports_image_generation);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn call_summary_is_always_stored_and_payloads_follow_detailed_setting() {
|
||||
let (_directory, store) = store().await;
|
||||
let provider = store
|
||||
.create_provider(&ProviderEndpointInput {
|
||||
name: "Local".into(),
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
base_url: "https://example.com/v1".into(),
|
||||
api_key: None,
|
||||
custom_headers: serde_json::json!({}),
|
||||
extra_params: serde_json::json!({}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let model = store
|
||||
.save_provider_model(
|
||||
provider.provider_id,
|
||||
&ProviderModelInput {
|
||||
model_id: "model-a".into(),
|
||||
display_name: "Model A".into(),
|
||||
endpoint_type: ProviderType::OpenAiChat,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: None,
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let call = NewLlmCall {
|
||||
call_id: "call-1".into(),
|
||||
run_id: "run-1".into(),
|
||||
conversation_id: "conversation-1".into(),
|
||||
provider_call_index: 0,
|
||||
model_hash: model.model_hash,
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
provider_url: provider.base_url,
|
||||
request_type: ProviderType::OpenAiChat,
|
||||
request_url: "https://example.com/v1/chat/completions".into(),
|
||||
model_id: model.model_id,
|
||||
display_name: model.display_name,
|
||||
reasoning_effort: Some("high".into()),
|
||||
fast: true,
|
||||
message_count: 2,
|
||||
tool_count: 3,
|
||||
detailed: false,
|
||||
};
|
||||
let conversation_id = ConversationId::new("conversation-1");
|
||||
let base_revision_id = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
store
|
||||
.claim_run(&PreparedRun {
|
||||
run_id: RunId::new("run-1"),
|
||||
conversation_id,
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new(call.model_hash.clone()),
|
||||
prompt: PromptSpec {
|
||||
instructions: String::new(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
initial_messages: Vec::new(),
|
||||
action: RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
},
|
||||
base_revision_id,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
store.start_llm_call(&call).await.unwrap();
|
||||
store
|
||||
.record_llm_request(
|
||||
"call-1",
|
||||
&serde_json::json!({}),
|
||||
&serde_json::json!({"model":"model-a"}),
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.record_llm_chunk("call-1", 0, 4, b"data", false)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.record_llm_usage(
|
||||
"call-1",
|
||||
Usage {
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(5),
|
||||
total_tokens: Some(15),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.finish_llm_call("call-1", "completed", Some("stop"), 9, None, None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let summary = store.llm_call("call-1").await.unwrap().unwrap();
|
||||
assert_eq!(summary.total_tokens, Some(15));
|
||||
assert_eq!(summary.reasoning_effort.as_deref(), Some("high"));
|
||||
assert_eq!(summary.fast, Some(true));
|
||||
assert_eq!(summary.request_bytes, Some(19));
|
||||
assert_eq!(summary.response_bytes, 4);
|
||||
assert!(store.llm_call_request("call-1").await.unwrap().is_none());
|
||||
assert!(store.llm_call_chunks("call-1").await.unwrap().is_empty());
|
||||
|
||||
let abandoned = NewLlmCall {
|
||||
call_id: "call-2".into(),
|
||||
provider_call_index: 1,
|
||||
..call
|
||||
};
|
||||
store.start_llm_call(&abandoned).await.unwrap();
|
||||
store
|
||||
.finish_run(&RunId::new("run-1"), RunStatus::Cancelled, None, None)
|
||||
.await
|
||||
.unwrap();
|
||||
let abandoned = store.llm_call("call-2").await.unwrap().unwrap();
|
||||
assert_eq!(abandoned.status, "cancelled");
|
||||
assert!(abandoned.finished_at_ms.is_some());
|
||||
assert!(abandoned.duration_ms.is_some());
|
||||
assert_eq!(
|
||||
store.llm_call("call-1").await.unwrap().unwrap().status,
|
||||
"completed"
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,928 @@
|
||||
use cursor_server::{
|
||||
client::ClientEvent,
|
||||
config::{ProviderConfig, ProviderKind},
|
||||
model::{
|
||||
ContentPart, ModelInvocation, ModelLatency, ModelRequest, ModelSpec, ProjectedContent,
|
||||
ProjectedMessage, PromptSpec, Role, Usage,
|
||||
},
|
||||
provider::{
|
||||
FinishReason, ModelEvent, OpenAiChatProvider, OpenAiResponsesProvider, Provider,
|
||||
ProviderStream,
|
||||
},
|
||||
run::{consume_model_cycle, RunFailure},
|
||||
};
|
||||
use futures_util::{stream, StreamExt};
|
||||
use serde_json::Value;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
fn provider_stream(events: Vec<ModelEvent>) -> ProviderStream {
|
||||
Box::pin(stream::iter(events.into_iter().map(Ok)))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn complete_tool_stream_is_validated_and_projected() {
|
||||
let (sender, mut receiver) = tokio::sync::mpsc::channel(32);
|
||||
let result = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ThinkingStart,
|
||||
ModelEvent::ThinkingDelta("why".into()),
|
||||
ModelEvent::ThinkingEnd,
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: r#"{"path":"/tmp/a"}"#.into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(2),
|
||||
..Usage::default()
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
drop(sender);
|
||||
|
||||
assert_eq!(result.reasoning, "why");
|
||||
assert_eq!(result.calls[0].arguments["path"], "/tmp/a");
|
||||
assert_eq!(result.usage.unwrap().input_tokens, Some(10));
|
||||
let mut events = Vec::new();
|
||||
while let Some(event) = receiver.recv().await {
|
||||
events.push(event);
|
||||
}
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ClientEvent::ThinkingEnd { .. })));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn eof_and_half_a_tool_call_are_failures_and_keep_only_diagnostics() {
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(8);
|
||||
let failure = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{".into(),
|
||||
},
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(failure.failure, RunFailure::Provider(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn done_with_open_blocks_and_events_after_done_are_rejected() {
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(8);
|
||||
let open = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(open.failure, RunFailure::Protocol(_)));
|
||||
|
||||
let after = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
ModelEvent::Usage(Usage::default()),
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(after.failure, RunFailure::Protocol(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn duplicate_usage_is_rejected_instead_of_guessing_which_total_is_final() {
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(8);
|
||||
let failure = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::Usage(Usage::default()),
|
||||
ModelEvent::Usage(Usage::default()),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(failure.failure, RunFailure::Protocol(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_raw_stream_and_request_projection_match_the_endpoint() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
"/v1/chat/completions",
|
||||
concat!(
|
||||
"data: {\"choices\":[{\"delta\":{\"reasoning_content\":\"wh\"},\"finish_reason\":null}]}\n\n",
|
||||
"data: {\"choices\":[{\"delta\":{\"reasoning_content\":\"y\"},\"finish_reason\":null}]}\n\n",
|
||||
"data: {\"choices\":[{\"delta\":{\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}\n\n",
|
||||
"data: {\"choices\":[],\"usage\":{\"prompt_tokens\":7,\"completion_tokens\":2}}\n\n",
|
||||
"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiChatProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiChat, base_url, None),
|
||||
);
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
let body = requests.recv().await.unwrap();
|
||||
let continued = continued_invocation(&events);
|
||||
let _ = collect(provider.stream(continued, CancellationToken::new())).await;
|
||||
let second_body = requests.recv().await.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert_eq!(body["messages"][0]["content"], "system");
|
||||
assert_eq!(body["messages"][1]["content"][0]["text"], "hello");
|
||||
assert_eq!(body["messages"][1]["content"][1]["type"], "image_url");
|
||||
assert_eq!(
|
||||
body["messages"][1]["content"][1]["image_url"]["url"],
|
||||
"data:image/png;base64,AQID"
|
||||
);
|
||||
assert!(body.get("max_completion_tokens").is_none());
|
||||
assert!(body.get("reasoning_effort").is_none());
|
||||
assert!(body.get("service_tier").is_none());
|
||||
assert_eq!(
|
||||
events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
ModelEvent::ThinkingDelta(text) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>(),
|
||||
"why"
|
||||
);
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::Usage(usage)
|
||||
if usage.input_tokens == Some(7)
|
||||
&& usage.output_tokens == Some(2)
|
||||
&& usage.cache_read_tokens.is_none()
|
||||
&& usage.cache_write_tokens.is_none()
|
||||
&& usage.reasoning_tokens.is_none())));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
assert_array_prefix(&body["messages"], &second_body["messages"]);
|
||||
assert_eq!(
|
||||
second_body["messages"].as_array().unwrap().last().unwrap()["reasoning_content"],
|
||||
"why"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_done_marker_can_terminate_without_finish_reason() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/chat/completions",
|
||||
concat!(
|
||||
"data: {\"choices\":[{\"delta\":{\"content\":\"ok\"},\"finish_reason\":null}]}\n\n",
|
||||
"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiChatProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiChat, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_accepts_content_filter_finish_reason() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/chat/completions",
|
||||
"data: {\"choices\":[{\"delta\":{\"content\":\"partial\"},\"finish_reason\":\"content_filter\"}]}\n\n",
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiChatProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiChat, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_buffers_split_tool_metadata() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/chat/completions",
|
||||
concat!(
|
||||
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"call-1\",\"function\":{\"arguments\":\"{\\\"\"}}]},\"finish_reason\":null}]}\n\n",
|
||||
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"name\":\"Read\",\"arguments\":\"path\\\":\\\"a\\\"}\"}}]},\"finish_reason\":\"tool_calls\"}]}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiChatProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiChat, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events.iter().any(|event| matches!(event, ModelEvent::ToolCallStart { call_id, name, .. } if call_id == "call-1" && name == "Read")));
|
||||
assert_eq!(
|
||||
events.last(),
|
||||
Some(&ModelEvent::Done(FinishReason::ToolUse))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_raw_stream_does_not_invent_reasoning_effort() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"why\"}\n\n",
|
||||
"data: {\"type\":\"response.reasoning_summary_text.done\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"id\":\"r1\",\"encrypted_content\":\"opaque-1\"}}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"id\":\"r2\",\"encrypted_content\":\"opaque-2\"}}\n\n",
|
||||
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n",
|
||||
"data: {\"type\":\"response.output_text.done\"}\n\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":8,\"output_tokens\":3}}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let mut request = invocation();
|
||||
request.request.model.reasoning.enabled = true;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, Some(4096)),
|
||||
);
|
||||
let events = collect(provider.stream(request, CancellationToken::new())).await;
|
||||
let body = requests.recv().await.unwrap();
|
||||
let continued = continued_invocation(&events);
|
||||
let _ = collect(provider.stream(continued, CancellationToken::new())).await;
|
||||
let second_body = requests.recv().await.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert_eq!(body["reasoning"]["summary"], "auto");
|
||||
assert!(body["reasoning"].get("effort").is_none());
|
||||
assert!(body.get("service_tier").is_none());
|
||||
assert_eq!(body["max_output_tokens"], 4096);
|
||||
assert_eq!(body["input"][0]["content"][1]["type"], "input_image");
|
||||
assert_eq!(
|
||||
body["input"][0]["content"][1]["image_url"],
|
||||
"data:image/png;base64,AQID"
|
||||
);
|
||||
assert!(events.iter().any(|event| matches!(event, ModelEvent::ProviderReplayState(state) if state.provider_kind == "openai_responses")));
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::Usage(usage)
|
||||
if usage.input_tokens == Some(8)
|
||||
&& usage.output_tokens == Some(3)
|
||||
&& usage.cache_read_tokens.is_none()
|
||||
&& usage.cache_write_tokens.is_none()
|
||||
&& usage.reasoning_tokens.is_none())));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
assert_array_prefix(&body["input"], &second_body["input"]);
|
||||
let replayed = second_body["input"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|item| item.get("encrypted_content").and_then(Value::as_str))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(replayed, ["opaque-1", "opaque-2"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_reasoning_item_done_closes_an_open_summary() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"why\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"id\":\"r1\",\"encrypted_content\":\"opaque\"}}\n\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::ThinkingEnd)));
|
||||
assert!(events.iter().any(
|
||||
|event| matches!(event, ModelEvent::ProviderReplayState(state)
|
||||
if state.value["items"][0]["encrypted_content"] == "opaque")
|
||||
));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_item_done_closes_text_and_tool_arguments() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"ok\"}]}}\n\n",
|
||||
"data: {\"type\":\"response.output_item.added\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\"}}\n\n",
|
||||
"data: {\"type\":\"response.function_call_arguments.delta\",\"output_index\":1,\"delta\":\"{\\\"path\\\":\\\"a\\\"}\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}}\n\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::TextEnd)));
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::ToolCallEnd { index: 1 })));
|
||||
assert_eq!(
|
||||
events.last(),
|
||||
Some(&ModelEvent::Done(FinishReason::ToolUse))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_completed_object_recovers_missing_item_events() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"output\":[",
|
||||
"{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"ok\"}]},",
|
||||
"{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}",
|
||||
"]}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::TextDelta(text) if text == "ok")));
|
||||
assert!(events.iter().any(|event| matches!(event, ModelEvent::ToolCallStart { call_id, name, .. } if call_id == "call-1" && name == "Read")));
|
||||
assert_eq!(
|
||||
events.last(),
|
||||
Some(&ModelEvent::Done(FinishReason::ToolUse))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_done_marker_accepts_completed_items() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"ok\"}]}}\n\n",
|
||||
"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_fast_is_projected_as_service_tier_fast() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
"/v1/chat/completions",
|
||||
concat!(
|
||||
"data: {\"choices\":[{\"delta\":{\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}\n\n",
|
||||
"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiChatProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiChat, base_url, None),
|
||||
);
|
||||
let mut request = invocation();
|
||||
request.request.model.model_id = "GPT-5.6-sol".into();
|
||||
request.request.model.latency = ModelLatency::Fast;
|
||||
|
||||
let _ = collect(provider.stream(request, CancellationToken::new())).await;
|
||||
let body = requests.recv().await.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert_eq!(body["service_tier"], "fast");
|
||||
assert_eq!(body["prompt_cache_key"], "cursor-byok");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_fast_is_projected_as_service_tier_fast() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n",
|
||||
"data: {\"type\":\"response.output_text.done\"}\n\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = OpenAiResponsesProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
let mut request = invocation();
|
||||
request.request.model.model_id = "gpt-5.6-sol".into();
|
||||
request.request.model.latency = ModelLatency::Fast;
|
||||
|
||||
let _ = collect(provider.stream(request, CancellationToken::new())).await;
|
||||
let body = requests.recv().await.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert_eq!(body["service_tier"], "fast");
|
||||
assert_eq!(body["prompt_cache_key"], "cursor-byok");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_raw_stream_uses_explicit_and_default_token_limits() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"event: message_start\ndata: {\"message\":{\"usage\":{\"input_tokens\":9}}}\n\n",
|
||||
"event: content_block_start\ndata: {\"index\":0,\"content_block\":{\"type\":\"thinking\",\"thinking\":\"\",\"signature\":\"\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":0,\"delta\":{\"type\":\"thinking_delta\",\"thinking\":\"first\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":0,\"delta\":{\"type\":\"signature_delta\",\"signature\":\"signature-1\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":0}\n\n",
|
||||
"event: content_block_start\ndata: {\"index\":1,\"content_block\":{\"type\":\"thinking\",\"thinking\":\"\",\"signature\":\"\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":1,\"delta\":{\"type\":\"thinking_delta\",\"thinking\":\"second\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":1,\"delta\":{\"type\":\"signature_delta\",\"signature\":\"signature-2\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":1}\n\n",
|
||||
"event: content_block_start\ndata: {\"index\":2,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":2,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":2}\n\n",
|
||||
"event: message_delta\ndata: {\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":2}}\n\n",
|
||||
"event: message_stop\ndata: {}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url.clone(), Some(1234)),
|
||||
);
|
||||
let mut request = invocation();
|
||||
request.request.model.latency = ModelLatency::Fast;
|
||||
let events = collect(provider.stream(request, CancellationToken::new())).await;
|
||||
let body = requests.recv().await.unwrap();
|
||||
let continued = continued_invocation(&events);
|
||||
let _ = collect(provider.stream(continued, CancellationToken::new())).await;
|
||||
let second_body = requests.recv().await.unwrap();
|
||||
let default_provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
let _ = collect(default_provider.stream(invocation(), CancellationToken::new())).await;
|
||||
let default_body = requests.recv().await.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert_eq!(body["max_tokens"], 1234);
|
||||
assert_eq!(default_body["max_tokens"], 65_000);
|
||||
assert!(body.get("service_tier").is_none());
|
||||
assert_eq!(body["messages"][0]["content"][1]["type"], "image");
|
||||
assert_eq!(
|
||||
body["messages"][0]["content"][1]["source"]["media_type"],
|
||||
"image/png"
|
||||
);
|
||||
assert_eq!(body["messages"][0]["content"][1]["source"]["data"], "AQID");
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::Usage(usage)
|
||||
if usage.input_tokens == Some(9)
|
||||
&& usage.output_tokens == Some(2)
|
||||
&& usage.cache_read_tokens.is_none()
|
||||
&& usage.cache_write_tokens.is_none()
|
||||
&& usage.reasoning_tokens.is_none())));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
assert_array_prefix(&body["messages"], &second_body["messages"]);
|
||||
let signatures = second_body["messages"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.flat_map(|message| message["content"].as_array().into_iter().flatten())
|
||||
.filter_map(|block| block.get("signature").and_then(Value::as_str))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(signatures, ["signature-1", "signature-2"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_unsigned_thinking_is_displayed_but_not_replayed() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"event: content_block_start\ndata: {\"index\":0,\"content_block\":{\"type\":\"thinking\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":0,\"delta\":{\"type\":\"thinking_delta\",\"thinking\":\"why\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":0}\n\n",
|
||||
"event: message_delta\ndata: {\"delta\":{\"stop_reason\":\"end_turn\"}}\n\n",
|
||||
"event: message_stop\ndata: {}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::ThinkingDelta(text) if text == "why")));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::ProviderReplayState(_))));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_redacted_thinking_is_preserved_for_replay() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"event: content_block_start\ndata: {\"index\":0,\"content_block\":{\"type\":\"redacted_thinking\",\"data\":\"opaque\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":0}\n\n",
|
||||
"event: message_delta\ndata: {\"delta\":{\"stop_reason\":\"end_turn\"}}\n\n",
|
||||
"event: message_stop\ndata: {}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events.iter().any(
|
||||
|event| matches!(event, ModelEvent::ProviderReplayState(state)
|
||||
if state.value["blocks"][0]["type"] == "redacted_thinking"
|
||||
&& state.value["blocks"][0]["data"] == "opaque")
|
||||
));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_message_stop_closes_blocks_and_infers_finish_reason() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"event: content_block_start\ndata: {\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\n\n",
|
||||
"event: message_stop\ndata: {}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::TextEnd)));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_final_message_delta_can_terminate_without_message_stop() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"event: content_block_start\ndata: {\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\n",
|
||||
"event: content_block_delta\ndata: {\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\n\n",
|
||||
"event: content_block_stop\ndata: {\"index\":0}\n\n",
|
||||
"event: message_delta\ndata: {\"delta\":{\"stop_reason\":\"refusal\"}}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_accepts_event_type_from_data() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/messages",
|
||||
concat!(
|
||||
"data: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\"}}\n\n",
|
||||
"data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\n\n",
|
||||
"data: {\"type\":\"content_block_stop\",\"index\":0}\n\n",
|
||||
"data: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\n",
|
||||
"data: {\"type\":\"message_stop\"}\n\n",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let provider = cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::Anthropic, base_url, None),
|
||||
);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::TextDelta(text) if text == "ok")));
|
||||
assert_eq!(events.last(), Some(&ModelEvent::Done(FinishReason::Stop)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_tool_arguments_are_normalized_but_nonempty_invalid_json_is_rejected() {
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(16);
|
||||
let empty = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
name: "NoArgs".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(empty.calls[0].arguments, serde_json::json!({}));
|
||||
|
||||
let invalid = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
name: "Broken".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: "{".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(invalid.failure, RunFailure::Protocol(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn every_provider_can_be_cancelled_while_waiting_for_response_headers() {
|
||||
for kind in [
|
||||
ProviderKind::OpenAiChat,
|
||||
ProviderKind::OpenAiResponses,
|
||||
ProviderKind::Anthropic,
|
||||
] {
|
||||
let (base_url, accepted, server) = hanging_server().await;
|
||||
let config = config(kind.clone(), base_url, Some(1024));
|
||||
let provider: Arc<dyn Provider> = match kind {
|
||||
ProviderKind::OpenAiChat => {
|
||||
Arc::new(OpenAiChatProvider::new(reqwest::Client::new(), config))
|
||||
}
|
||||
ProviderKind::OpenAiResponses => {
|
||||
Arc::new(OpenAiResponsesProvider::new(reqwest::Client::new(), config))
|
||||
}
|
||||
ProviderKind::Anthropic => Arc::new(cursor_server::provider::AnthropicProvider::new(
|
||||
reqwest::Client::new(),
|
||||
config,
|
||||
)),
|
||||
};
|
||||
let cancellation = CancellationToken::new();
|
||||
let stream = provider.stream(invocation(), cancellation.clone());
|
||||
let collect = tokio::spawn(async move { collect(stream).await });
|
||||
accepted.await.unwrap();
|
||||
cancellation.cancel();
|
||||
let events = tokio::time::timeout(Duration::from_secs(1), collect)
|
||||
.await
|
||||
.expect("provider did not cancel while waiting for headers")
|
||||
.unwrap();
|
||||
assert!(events.is_empty());
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
|
||||
fn invocation() -> ModelInvocation {
|
||||
ModelInvocation {
|
||||
call_id: "call-1".into(),
|
||||
run_id: "run-1".into(),
|
||||
conversation_id: "conversation-1".into(),
|
||||
provider_call_index: 0,
|
||||
request: ModelRequest {
|
||||
prompt: PromptSpec {
|
||||
instructions: "system".into(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
model: ModelSpec::new("model"),
|
||||
history: vec![ProjectedMessage {
|
||||
message_id: "user-1".into(),
|
||||
role: Role::User,
|
||||
content: ProjectedContent::Parts(vec![
|
||||
ContentPart::Text {
|
||||
text: "hello".into(),
|
||||
},
|
||||
ContentPart::Image {
|
||||
mime_type: "image/png".into(),
|
||||
data: vec![1, 2, 3],
|
||||
},
|
||||
]),
|
||||
}],
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn continued_invocation(events: &[ModelEvent]) -> ModelInvocation {
|
||||
let mut invocation = invocation();
|
||||
let text = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
ModelEvent::TextDelta(text) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
let thinking = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
ModelEvent::ThinkingDelta(text) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<String>();
|
||||
let replay_state = events.iter().find_map(|event| match event {
|
||||
ModelEvent::ProviderReplayState(state) => Some(state.clone()),
|
||||
_ => None,
|
||||
});
|
||||
invocation.request.history.push(ProjectedMessage {
|
||||
message_id: "assistant-1".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text,
|
||||
thinking,
|
||||
replay_state,
|
||||
calls: Vec::new(),
|
||||
},
|
||||
});
|
||||
invocation.call_id = "call-2".into();
|
||||
invocation
|
||||
}
|
||||
|
||||
fn assert_array_prefix(first: &Value, second: &Value) {
|
||||
let first = first.as_array().unwrap();
|
||||
let second = second.as_array().unwrap();
|
||||
assert_eq!(first.as_slice(), &second[..first.len()]);
|
||||
}
|
||||
|
||||
fn config(kind: ProviderKind, base_url: String, max_output_tokens: Option<u64>) -> ProviderConfig {
|
||||
let path = match kind {
|
||||
ProviderKind::OpenAiChat => "/chat/completions",
|
||||
ProviderKind::OpenAiResponses => "/responses",
|
||||
ProviderKind::Anthropic => "/messages",
|
||||
};
|
||||
ProviderConfig {
|
||||
kind,
|
||||
request_url: format!("{base_url}{path}"),
|
||||
api_key: "test".into(),
|
||||
custom_headers: Default::default(),
|
||||
max_output_tokens,
|
||||
request_timeout: Duration::from_secs(5),
|
||||
}
|
||||
}
|
||||
|
||||
async fn collect(stream: ProviderStream) -> Vec<ModelEvent> {
|
||||
stream.map(|event| event.unwrap()).collect::<Vec<_>>().await
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct FixtureState {
|
||||
response: Arc<str>,
|
||||
requests: tokio::sync::mpsc::UnboundedSender<Value>,
|
||||
}
|
||||
|
||||
async fn fixture_server(
|
||||
path: &'static str,
|
||||
response: &'static str,
|
||||
) -> (
|
||||
String,
|
||||
tokio::sync::mpsc::UnboundedReceiver<Value>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
) {
|
||||
async fn endpoint(
|
||||
axum::extract::State(state): axum::extract::State<FixtureState>,
|
||||
axum::Json(body): axum::Json<Value>,
|
||||
) -> impl axum::response::IntoResponse {
|
||||
let _ = state.requests.send(body);
|
||||
(
|
||||
[(axum::http::header::CONTENT_TYPE, "text/event-stream")],
|
||||
state.response.to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
let (sender, receiver) = tokio::sync::mpsc::unbounded_channel();
|
||||
let app = axum::Router::new()
|
||||
.route(path, axum::routing::post(endpoint))
|
||||
.with_state(FixtureState {
|
||||
response: response.into(),
|
||||
requests: sender,
|
||||
});
|
||||
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, app).await.unwrap();
|
||||
});
|
||||
(format!("http://{address}/v1"), receiver, server)
|
||||
}
|
||||
|
||||
async fn hanging_server() -> (
|
||||
String,
|
||||
tokio::sync::oneshot::Receiver<()>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
) {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let (accepted, receiver) = tokio::sync::oneshot::channel();
|
||||
let server = tokio::spawn(async move {
|
||||
let (_socket, _) = listener.accept().await.unwrap();
|
||||
let _ = accepted.send(());
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
(format!("http://{address}/v1"), receiver, server)
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use cursor_server::model::{
|
||||
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind,
|
||||
};
|
||||
|
||||
fn prepared(
|
||||
run_id: &str,
|
||||
conversation_id: &ConversationId,
|
||||
base_revision_id: cursor_server::model::RevisionId,
|
||||
) -> PreparedRun {
|
||||
PreparedRun {
|
||||
run_id: RunId::new(run_id),
|
||||
conversation_id: conversation_id.clone(),
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new("test-model"),
|
||||
prompt: PromptSpec {
|
||||
instructions: "test".into(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
initial_messages: Vec::new(),
|
||||
action: RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
},
|
||||
base_revision_id,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selecting_an_old_revision_creates_a_branch_without_old_suffixes() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("conversation");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let first = prepared("run-1", &conversation_id, root);
|
||||
store.claim_run(&first).await.unwrap();
|
||||
|
||||
let a = fixtures::user("a", "A");
|
||||
let revision_a = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&first.run_id,
|
||||
root,
|
||||
std::slice::from_ref(&a),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let b = fixtures::user("b", "B");
|
||||
let revision_b = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&first.run_id,
|
||||
revision_a,
|
||||
std::slice::from_ref(&b),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let second = prepared("run-2", &conversation_id, revision_a);
|
||||
let claimed = store.claim_run(&second).await.unwrap();
|
||||
assert_eq!(claimed.replaced_run_id.as_ref(), Some(&first.run_id));
|
||||
let c = fixtures::user("c", "C");
|
||||
let revision_c = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&second.run_id,
|
||||
revision_a,
|
||||
std::slice::from_ref(&c),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
store.load_revision_messages(revision_b).await.unwrap(),
|
||||
vec![a.clone(), b]
|
||||
);
|
||||
assert_eq!(
|
||||
store.load_revision_messages(revision_c).await.unwrap(),
|
||||
vec![a, c]
|
||||
);
|
||||
assert!(store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&first.run_id,
|
||||
revision_b,
|
||||
&[fixtures::user("late", "late")],
|
||||
)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn identical_runtime_event_is_exactly_once_and_conflicts_are_rejected() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("runtime");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let run = prepared("run", &conversation_id, root);
|
||||
store.claim_run(&run).await.unwrap();
|
||||
let event = cursor_server::model::RuntimeEvent {
|
||||
event_id: "branch:changed:7".into(),
|
||||
text: "runtime state changed".into(),
|
||||
}
|
||||
.into_message();
|
||||
let (revision, inserted) = store
|
||||
.append_message_once(&conversation_id, &run.run_id, root, &event)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(inserted);
|
||||
let (same, inserted) = store
|
||||
.append_message_once(&conversation_id, &run.run_id, revision, &event)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(same, revision);
|
||||
assert!(!inserted);
|
||||
|
||||
let conflict = cursor_server::model::RuntimeEvent {
|
||||
event_id: "branch:changed:7".into(),
|
||||
text: "different".into(),
|
||||
}
|
||||
.into_message();
|
||||
assert!(store
|
||||
.append_message_once(&conversation_id, &run.run_id, revision, &conflict)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn editing_a_logical_input_discards_its_active_suffix() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("edited-conversation");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let input_id = "cursor:user:stable-id";
|
||||
assert_eq!(
|
||||
store
|
||||
.anchor_input(&conversation_id, input_id, root)
|
||||
.await
|
||||
.unwrap(),
|
||||
root
|
||||
);
|
||||
|
||||
let first = prepared("first-run", &conversation_id, root);
|
||||
store.claim_run(&first).await.unwrap();
|
||||
let original = fixtures::user("original", "original text");
|
||||
let original_revision = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&first.run_id,
|
||||
root,
|
||||
std::slice::from_ref(&original),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let suffix = fixtures::user("suffix", "old suffix");
|
||||
let old_head = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&first.run_id,
|
||||
original_revision,
|
||||
std::slice::from_ref(&suffix),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let edit_base = store
|
||||
.anchor_input(&conversation_id, input_id, old_head)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(edit_base, root);
|
||||
let second = prepared("second-run", &conversation_id, edit_base);
|
||||
store.claim_run(&second).await.unwrap();
|
||||
let edited = fixtures::user("edited", "edited text");
|
||||
let edited_head = store
|
||||
.append_revision(
|
||||
&conversation_id,
|
||||
&second.run_id,
|
||||
edit_base,
|
||||
std::slice::from_ref(&edited),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
store.load_revision_messages(edited_head).await.unwrap(),
|
||||
vec![edited]
|
||||
);
|
||||
assert_eq!(
|
||||
store.load_revision_messages(old_head).await.unwrap(),
|
||||
vec![original, suffix]
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,475 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
CursorCommand, CursorSessionRegistry,
|
||||
},
|
||||
model::{ContentPart, ProjectedContent},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
store::{BlobId, Store},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn current_mode_and_referenced_context_are_consumed_by_one_runtime_message() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let references = references(&store).await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("answer".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("ask-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(run_request(references)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
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 message = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
}
|
||||
|
||||
let requests = provider.requests();
|
||||
let request = &requests[0];
|
||||
assert!(request
|
||||
.prompt
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "AskQuestion"));
|
||||
assert!(!request
|
||||
.prompt
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "GenerateImage"));
|
||||
assert_eq!(request.history.len(), 1);
|
||||
assert_eq!(
|
||||
request.history[0].message_id,
|
||||
"runtime:run-request:ask-request"
|
||||
);
|
||||
let ProjectedContent::Parts(parts) = &request.history[0].content else {
|
||||
panic!("runtime message must use typed parts")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("this fixture has no images")
|
||||
};
|
||||
for expected in [
|
||||
"<user_rule>\nworkspace rule\n</user_rule>",
|
||||
"<agent_skill fullPath=\"/skills/test/SKILL.md\">test skill</agent_skill>",
|
||||
"<subagent name=\"reviewer\">review code</subagent>",
|
||||
"<mcp_meta_tool_server name=\"test\" identifier=\"mcp-test\">",
|
||||
"<mcp_tool name=\"lookup\">",
|
||||
"<definition_path>/tmp/mcp-test/lookup.json</definition_path>",
|
||||
"<input_schema>{"properties":{"query":{"type":"string"}},"type":"object"}</input_schema>",
|
||||
"Call a listed tool directly with CallMcpTool without calling GetMcpTools first.",
|
||||
"Ask mode is active.",
|
||||
"<user_query>\nexplain this\n</user_query>",
|
||||
] {
|
||||
assert!(
|
||||
text.contains(expected),
|
||||
"missing runtime section: {expected}"
|
||||
);
|
||||
}
|
||||
assert!(!text.contains("complete skill body"));
|
||||
assert!(!text.contains("complete MCP server instructions"));
|
||||
assert!(text.contains("/workspace/src/main.rs"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn missing_context_parts_use_current_cursor_response_and_cache_its_content() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let referenced_context = fixture_context();
|
||||
let references = references_for(&referenced_context);
|
||||
let mut current_context = referenced_context.clone();
|
||||
current_context
|
||||
.mcp_meta_tool_options
|
||||
.as_mut()
|
||||
.unwrap()
|
||||
.mcp_descriptors
|
||||
.push(pb::McpDescriptor {
|
||||
server_identifier: "live-mcp".into(),
|
||||
tools: vec![pb::McpToolDescriptor {
|
||||
tool_name: "current-tool".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
});
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("answer".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("context-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(run_request(references)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let mut requested_context = 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 {
|
||||
break;
|
||||
}
|
||||
let message = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match message.message {
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
assert_eq!(exec.id, 0);
|
||||
let Some(pb::exec_server_message::Message::RequestContextArgs(args)) = exec.message
|
||||
else {
|
||||
panic!("missing context must use RequestContextArgs")
|
||||
};
|
||||
assert_eq!(args.notes_session_id.as_deref(), Some("mode-conversation"));
|
||||
requested_context = true;
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(context_stream_close()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: 0,
|
||||
message: Some(
|
||||
pb::exec_client_message::Message::RequestContextResult(
|
||||
pb::RequestContextResult {
|
||||
result: Some(
|
||||
pb::request_context_result::Result::Success(
|
||||
pb::RequestContextSuccess {
|
||||
request_context: Some(
|
||||
current_context.clone(),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
assert!(requested_context);
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 1);
|
||||
let ProjectedContent::Parts(parts) = &requests[0].history[0].content else {
|
||||
panic!("runtime message must use typed parts")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("this fixture has no images")
|
||||
};
|
||||
assert!(text.contains("<mcp_meta_tool_server name=\"live-mcp\" identifier=\"live-mcp\">"));
|
||||
assert!(text.contains("<mcp_tool name=\"current-tool\">"));
|
||||
|
||||
let stale = references_for(&referenced_context);
|
||||
let current = references_for(¤t_context);
|
||||
for (id, _) in [
|
||||
current.rules,
|
||||
current.skills,
|
||||
current.subagents,
|
||||
current.mcps,
|
||||
] {
|
||||
assert!(store.get_blob(&id).await.unwrap().is_some());
|
||||
}
|
||||
assert!(store.get_blob(&stale.mcps.0).await.unwrap().is_none());
|
||||
}
|
||||
|
||||
struct References {
|
||||
rules: (BlobId, u32),
|
||||
skills: (BlobId, u32),
|
||||
subagents: (BlobId, u32),
|
||||
mcps: (BlobId, u32),
|
||||
}
|
||||
|
||||
async fn references(store: &Store) -> References {
|
||||
let context = fixture_context();
|
||||
let references = references_for(&context);
|
||||
for data in part_data(&context) {
|
||||
store.put_blob(&data, &[]).await.unwrap();
|
||||
}
|
||||
references
|
||||
}
|
||||
|
||||
fn fixture_context() -> pb::RequestContext {
|
||||
pb::RequestContext {
|
||||
rules: vec![pb::CursorRule {
|
||||
full_path: "/workspace/AGENTS.md".into(),
|
||||
content: "workspace rule".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
non_file_rules: vec![pb::CursorRule {
|
||||
full_path: "/skills/test/SKILL.md".into(),
|
||||
content: "complete skill body".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
agent_skills: vec![pb::AgentSkill {
|
||||
full_path: "/skills/test/SKILL.md".into(),
|
||||
description: "test skill".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
custom_subagents: vec![pb::CustomSubagent {
|
||||
name: "reviewer".into(),
|
||||
description: "review code".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
mcp_meta_tool_options: Some(pb::McpMetaToolOptions {
|
||||
enabled: true,
|
||||
mcp_descriptors: vec![pb::McpDescriptor {
|
||||
server_name: "test".into(),
|
||||
server_identifier: "mcp-test".into(),
|
||||
server_use_instructions: Some("complete MCP server instructions".into()),
|
||||
tools: vec![pb::McpToolDescriptor {
|
||||
tool_name: "lookup".into(),
|
||||
definition_path: Some("/tmp/mcp-test/lookup.json".into()),
|
||||
description: Some("look up a value".into()),
|
||||
input_schema_json: Some(
|
||||
r#"{"type":"object","properties":{"query":{"type":"string"}}}"#.into(),
|
||||
),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
}],
|
||||
}),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
fn references_for(context: &pb::RequestContext) -> References {
|
||||
let mut parts = part_data(context).into_iter();
|
||||
References {
|
||||
rules: reference(&parts.next().unwrap()),
|
||||
skills: reference(&parts.next().unwrap()),
|
||||
subagents: reference(&parts.next().unwrap()),
|
||||
mcps: reference(&parts.next().unwrap()),
|
||||
}
|
||||
}
|
||||
|
||||
fn part_data(context: &pb::RequestContext) -> Vec<Vec<u8>> {
|
||||
vec![
|
||||
pb::RequestContextRulesPart {
|
||||
rules: context.rules.clone(),
|
||||
non_file_rules: context.non_file_rules.clone(),
|
||||
cloud_rule: context.cloud_rule.clone(),
|
||||
}
|
||||
.encode_to_vec(),
|
||||
pb::RequestContextSkillsPart {
|
||||
agent_skills: context.agent_skills.clone(),
|
||||
skill_options: context.skill_options.clone(),
|
||||
}
|
||||
.encode_to_vec(),
|
||||
pb::RequestContextSubagentsPart {
|
||||
custom_subagents: context.custom_subagents.clone(),
|
||||
}
|
||||
.encode_to_vec(),
|
||||
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(),
|
||||
}
|
||||
.encode_to_vec(),
|
||||
]
|
||||
}
|
||||
|
||||
fn reference(data: &[u8]) -> (BlobId, u32) {
|
||||
(BlobId::digest(data), data.len() as u32)
|
||||
}
|
||||
|
||||
fn run_request(references: References) -> pb::AgentClientMessage {
|
||||
let (rules, rules_byte_length) = references.rules;
|
||||
let (skills, skills_byte_length) = references.skills;
|
||||
let (subagents, subagents_byte_length) = references.subagents;
|
||||
let (mcps, mcps_byte_length) = references.mcps;
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
conversation_state: Some(pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
request_context_parts: Some(pb::RequestContextPartReferences {
|
||||
rules_blob_id: rules.as_bytes().to_vec(),
|
||||
rules_byte_length,
|
||||
skills_blob_id: skills.as_bytes().to_vec(),
|
||||
skills_byte_length,
|
||||
subagents_blob_id: subagents.as_bytes().to_vec(),
|
||||
subagents_byte_length,
|
||||
mcps_blob_id: mcps.as_bytes().to_vec(),
|
||||
mcps_byte_length,
|
||||
dynamic_context: Some(pb::RequestContext {
|
||||
env: Some(pb::RequestContextEnv {
|
||||
os_version: "darwin".into(),
|
||||
workspace_paths: vec!["/workspace".into()],
|
||||
shell: "zsh".into(),
|
||||
time_zone: "UTC".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
}),
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "explain this".into(),
|
||||
message_id: "wire-user".into(),
|
||||
mode: pb::AgentMode::Ask as i32,
|
||||
selected_context: Some(pb::SelectedContext {
|
||||
invocation_context: Some(pb::InvocationContext {
|
||||
data: Some(pb::invocation_context::Data::IdeState(
|
||||
pb::invocation_context::IdeState {
|
||||
visible_files: vec![
|
||||
pb::invocation_context::ide_state::File {
|
||||
path: "/workspace/src/main.rs".into(),
|
||||
total_lines: 10,
|
||||
..Default::default()
|
||||
},
|
||||
],
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("mode-conversation".into()),
|
||||
run_id: Some("wire-run".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn context_stream_close() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientControlMessage(
|
||||
pb::ExecClientControlMessage {
|
||||
message: Some(pb::exec_client_control_message::Message::StreamClose(
|
||||
pb::ExecClientStreamClose { id: 0 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use cursor_server::model::{
|
||||
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind, RuntimeEvent,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_event_is_appended_exactly_once() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("conversation");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let run = PreparedRun {
|
||||
run_id: RunId::new("run"),
|
||||
conversation_id: conversation_id.clone(),
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new("model"),
|
||||
prompt: PromptSpec {
|
||||
instructions: String::new(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
initial_messages: Vec::new(),
|
||||
action: RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
},
|
||||
base_revision_id: root,
|
||||
};
|
||||
store.claim_run(&run).await.unwrap();
|
||||
let message = RuntimeEvent {
|
||||
event_id: "branch:changed:7".into(),
|
||||
text: "runtime state changed".into(),
|
||||
}
|
||||
.into_message();
|
||||
|
||||
let (revision, inserted) = store
|
||||
.append_message_once(&conversation_id, &run.run_id, root, &message)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(inserted);
|
||||
assert!(
|
||||
!store
|
||||
.append_message_once(&conversation_id, &run.run_id, revision, &message)
|
||||
.await
|
||||
.unwrap()
|
||||
.1
|
||||
);
|
||||
let messages = store.load_revision_messages(revision).await.unwrap();
|
||||
assert_eq!(messages.len(), 1);
|
||||
assert_eq!(
|
||||
messages[0].runtime_event_id.as_deref(),
|
||||
Some("branch:changed:7")
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
CursorCommand, CursorSessionRegistry,
|
||||
},
|
||||
model::{ContentPart, ProjectedContent},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
store::BlobId,
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_checkpoint() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let stored_image = vec![4, 5];
|
||||
let stored_image_id = store.put_blob(&stored_image, &[]).await.unwrap();
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("seen".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("image-run").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(request(&stored_image_id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let mut checkpoints = Vec::new();
|
||||
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 message = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match message.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
checkpoints.push(state)
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
let requests = provider.requests();
|
||||
let user = requests[0]
|
||||
.history
|
||||
.iter()
|
||||
.find(|message| message.message_id == "runtime:run-request:image-run")
|
||||
.unwrap();
|
||||
let ProjectedContent::Parts(parts) = &user.content else {
|
||||
panic!("runtime user message must retain typed parts")
|
||||
};
|
||||
assert!(matches!(
|
||||
&parts[0],
|
||||
ContentPart::Text { text } if text.contains("<user_query>\nwhat is this?\n</user_query>")
|
||||
));
|
||||
assert_eq!(
|
||||
&parts[1..],
|
||||
&[
|
||||
ContentPart::Image {
|
||||
mime_type: "image/png".into(),
|
||||
data: vec![1, 2, 3],
|
||||
},
|
||||
ContentPart::Image {
|
||||
mime_type: "image/webp".into(),
|
||||
data: stored_image,
|
||||
},
|
||||
ContentPart::Image {
|
||||
mime_type: "image/jpeg".into(),
|
||||
data: vec![6, 7, 8],
|
||||
},
|
||||
]
|
||||
);
|
||||
let serialized_request = serde_json::to_string(&requests[0]).unwrap();
|
||||
assert!(!serialized_request.contains("private-image-uuid"));
|
||||
assert!(!serialized_request.contains("/private/image/path"));
|
||||
|
||||
let state = checkpoints
|
||||
.iter()
|
||||
.find(|state| state.root_prompt_messages_json.len() >= 2)
|
||||
.unwrap();
|
||||
let mut user_root = None;
|
||||
for raw_id in &state.root_prompt_messages_json {
|
||||
let id = BlobId::from_bytes(raw_id).unwrap();
|
||||
let bytes = store.get_blob(&id).await.unwrap().unwrap();
|
||||
let value: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
|
||||
if value["id"] == "runtime:run-request:image-run" {
|
||||
user_root = Some(value);
|
||||
break;
|
||||
}
|
||||
}
|
||||
let user_root = user_root.unwrap();
|
||||
assert_eq!(user_root["content"][1]["type"], "image");
|
||||
assert_eq!(user_root["content"][1]["mimeType"], "image/png");
|
||||
assert_eq!(user_root["content"][1]["image"], "AQID");
|
||||
assert!(user_root["content"][1].get("data").is_none());
|
||||
assert_eq!(user_root["content"][2]["mimeType"], "image/webp");
|
||||
assert_eq!(user_root["content"][2]["image"], "BAU=");
|
||||
assert!(user_root["content"][2].get("data").is_none());
|
||||
assert_eq!(user_root["content"][3]["mimeType"], "image/jpeg");
|
||||
assert_eq!(user_root["content"][3]["image"], "BgcI");
|
||||
assert!(user_root["content"][3].get("data").is_none());
|
||||
}
|
||||
|
||||
fn request(stored_image_id: &BlobId) -> pb::AgentClientMessage {
|
||||
let data = vec![1, 2, 3];
|
||||
let inline_blob_data = vec![6, 7, 8];
|
||||
let inline_blob_id = BlobId::digest(&inline_blob_data);
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "what is this?".into(),
|
||||
message_id: "image-user".into(),
|
||||
selected_context: Some(pb::SelectedContext {
|
||||
selected_images: vec![
|
||||
pb::SelectedImage {
|
||||
uuid: "private-image-uuid".into(),
|
||||
path: "/private/image/path".into(),
|
||||
mime_type: "image/png".into(),
|
||||
data_or_blob_id: Some(
|
||||
pb::selected_image::DataOrBlobId::Data(data),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
pb::SelectedImage {
|
||||
uuid: "stored-image".into(),
|
||||
path: "/private/stored/path".into(),
|
||||
mime_type: "image/webp".into(),
|
||||
data_or_blob_id: Some(
|
||||
pb::selected_image::DataOrBlobId::BlobId(
|
||||
stored_image_id.as_bytes().to_vec(),
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
pb::SelectedImage {
|
||||
uuid: "inline-blob-image".into(),
|
||||
path: "/private/inline/path".into(),
|
||||
mime_type: "image/jpeg".into(),
|
||||
data_or_blob_id: Some(
|
||||
pb::selected_image::DataOrBlobId::BlobIdWithData(
|
||||
pb::selected_image::BlobIdWithData {
|
||||
blob_id: inline_blob_id.as_bytes().to_vec(),
|
||||
data: inline_blob_data,
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
],
|
||||
..Default::default()
|
||||
}),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("image-conversation".into()),
|
||||
run_id: Some("image-run".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,299 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
connect,
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
proto::agent::v1 as pb,
|
||||
CursorCommand, CursorSessionHandle, CursorSessionRegistry,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn every_bidi_run_resolves_and_persists_its_own_subagent_model_and_background_state() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
for suffix in ["a", "b"] {
|
||||
provider.push(task_response(suffix));
|
||||
provider.push(stop_response(suffix));
|
||||
}
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
|
||||
let first = registry.get_or_create("subagent-run-a").await.unwrap();
|
||||
let first_checkpoint = drive(
|
||||
&first,
|
||||
run_request("subagent-run-a", "user-a", "model-a", None),
|
||||
"model-a",
|
||||
"child-a",
|
||||
)
|
||||
.await;
|
||||
|
||||
let second = registry.get_or_create("subagent-run-b").await.unwrap();
|
||||
drive(
|
||||
&second,
|
||||
run_request(
|
||||
"subagent-run-b",
|
||||
"user-b",
|
||||
"model-b",
|
||||
Some(first_checkpoint),
|
||||
),
|
||||
"model-b",
|
||||
"child-b",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
async fn drive(
|
||||
handle: &CursorSessionHandle,
|
||||
request: pb::AgentClientMessage,
|
||||
expected_model: &str,
|
||||
child_id: &str,
|
||||
) -> pb::ConversationStateStructure {
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(request),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let mut saw_started = false;
|
||||
let mut saw_exec = false;
|
||||
let mut saw_completed = false;
|
||||
let mut checkpoint = None;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(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;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
let Some(pb::exec_server_message::Message::SubagentArgs(args)) = exec.message
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
assert_eq!(args.model_id, expected_model);
|
||||
assert_eq!(args.run_in_background, Some(true));
|
||||
saw_exec = true;
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(subagent_result(exec.id, child_id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
match update.message {
|
||||
Some(pb::interaction_update::Message::ToolCallStarted(started)) => {
|
||||
let task = task(started.tool_call.as_ref().unwrap());
|
||||
assert_eq!(
|
||||
task.args.as_ref().unwrap().model.as_deref(),
|
||||
Some(expected_model)
|
||||
);
|
||||
saw_started = true;
|
||||
}
|
||||
Some(pb::interaction_update::Message::ToolCallCompleted(completed)) => {
|
||||
let task = task(completed.tool_call.as_ref().unwrap());
|
||||
let Some(pb::task_result::Result::Success(success)) = task
|
||||
.result
|
||||
.as_ref()
|
||||
.and_then(|result| result.result.as_ref())
|
||||
else {
|
||||
panic!("expected Task success")
|
||||
};
|
||||
assert_eq!(
|
||||
task.args.as_ref().unwrap().model.as_deref(),
|
||||
Some(expected_model)
|
||||
);
|
||||
assert!(success.is_background);
|
||||
assert_eq!(success.agent_id.as_deref(), Some(child_id));
|
||||
saw_completed = true;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
|
||||
if state.pending_tool_calls.is_empty() =>
|
||||
{
|
||||
checkpoint = Some(state);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
assert!(saw_started && saw_exec && saw_completed);
|
||||
let checkpoint = checkpoint.expect("settled checkpoint");
|
||||
let state = checkpoint
|
||||
.subagent_states
|
||||
.get(child_id)
|
||||
.expect("background subagent persisted state");
|
||||
assert_eq!(state.model_id.as_deref(), Some(expected_model));
|
||||
let run = checkpoint
|
||||
.subagent_runs_by_parent_tool_call_id
|
||||
.get(&format!("task-{child_id}"))
|
||||
.expect("background subagent run state");
|
||||
assert_eq!(run.status, pb::SubagentRunStatus::Backgrounded as i32);
|
||||
checkpoint
|
||||
}
|
||||
|
||||
fn task(call: &pb::ToolCall) -> &pb::TaskToolCall {
|
||||
let Some(pb::tool_call::Tool::TaskToolCall(task)) = call.tool.as_ref() else {
|
||||
panic!("expected TaskToolCall")
|
||||
};
|
||||
task
|
||||
}
|
||||
|
||||
fn run_request(
|
||||
request_id: &str,
|
||||
user_id: &str,
|
||||
subagent_model: &str,
|
||||
conversation_state: Option<pb::ConversationStateStructure>,
|
||||
) -> 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: format!("start {request_id}"),
|
||||
message_id: user_id.into(),
|
||||
mode: pb::AgentMode::Multitask as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("subagent-e2e-conversation".into()),
|
||||
run_id: Some(request_id.into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "parent-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state,
|
||||
subagent_model_overrides: vec![pb::SubagentModelOverride {
|
||||
subagent_type: "generalPurpose".into(),
|
||||
selection: Some(pb::subagent_model_override::Selection::Model(
|
||||
pb::RequestedModel {
|
||||
model_id: subagent_model.into(),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}],
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn task_response(suffix: &str) -> Vec<ModelEvent> {
|
||||
let child_id = format!("child-{suffix}");
|
||||
let arguments = serde_json::json!({
|
||||
"description": format!("background {suffix}"),
|
||||
"prompt": "inspect",
|
||||
"subagent_type": "generalPurpose",
|
||||
"run_in_background": true
|
||||
})
|
||||
.to_string();
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: format!("model-call-{suffix}"),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: format!("task-{child_id}"),
|
||||
name: "Task".into(),
|
||||
},
|
||||
ModelEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
delta: arguments,
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::Done(FinishReason::ToolUse),
|
||||
]
|
||||
}
|
||||
|
||||
fn stop_response(suffix: &str) -> Vec<ModelEvent> {
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: format!("final-{suffix}"),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("background task started".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]
|
||||
}
|
||||
|
||||
fn subagent_result(id: u32, child_id: &str) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::SubagentResult(
|
||||
pb::SubagentResult {
|
||||
result: Some(pb::subagent_result::Result::Success(pb::SubagentSuccess {
|
||||
agent_id: child_id.into(),
|
||||
final_message: Some("running in background".into()),
|
||||
background_reason: pb::SubagentBackgroundReason::AgentRequest as i32,
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..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,213 @@
|
||||
use std::collections::{BTreeMap, HashMap, HashSet};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
interaction,
|
||||
proto::agent::v1 as pb,
|
||||
tools::{
|
||||
codec,
|
||||
runtime::{CursorToolRuntime, ExecContext, SubagentModel},
|
||||
ToolBatchState, ToolDispatcher,
|
||||
},
|
||||
},
|
||||
model::{CanonicalMessage, MessageContent, Origin, Role, ToolCall},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn task_keeps_wire_type_model_parent_and_background_fields() {
|
||||
let mut context = context();
|
||||
context.subagent_models.insert(
|
||||
"cursor-guide".into(),
|
||||
SubagentModel::Model("guide-model".into()),
|
||||
);
|
||||
let call = task_call(serde_json::json!({
|
||||
"description": "guide",
|
||||
"prompt": "inspect",
|
||||
"subagent_type": "cursor-guide",
|
||||
"run_in_background": true,
|
||||
"interrupt": false,
|
||||
}));
|
||||
|
||||
let call = context.prepare_call(&call).unwrap();
|
||||
let message = codec::request(7, &call, &context).unwrap();
|
||||
let pb::agent_server_message::Message::ExecServerMessage(exec) = message.message.unwrap()
|
||||
else {
|
||||
panic!("expected ExecServerMessage")
|
||||
};
|
||||
assert_eq!(exec.id, 7);
|
||||
assert_eq!(exec.exec_id, "task-call");
|
||||
assert_eq!(exec.accept_hook_additional_contexts, Some(false));
|
||||
let pb::exec_server_message::Message::SubagentArgs(args) = exec.message.unwrap() else {
|
||||
panic!("expected SubagentArgs")
|
||||
};
|
||||
assert_eq!(args.subagent_type, "cursor-guide");
|
||||
assert_eq!(args.model_id, "guide-model");
|
||||
assert_eq!(args.parent_conversation_id.as_deref(), Some("child"));
|
||||
assert_eq!(args.root_parent_conversation_id.as_deref(), Some("root"));
|
||||
assert_eq!(args.run_in_background, Some(true));
|
||||
assert_eq!(args.interrupt, Some(false));
|
||||
let rendered = interaction::render_tool_call(&call, false).unwrap();
|
||||
let Some(pb::tool_call::Tool::TaskToolCall(task)) = rendered.tool else {
|
||||
panic!("expected TaskToolCall")
|
||||
};
|
||||
assert_eq!(task.args.unwrap().model.as_deref(), Some("guide-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_uses_explicit_call_model_then_the_run_default() {
|
||||
let mut explicit = task_call(serde_json::json!({
|
||||
"description": "task",
|
||||
"prompt": "inspect",
|
||||
"subagent_type": "generalPurpose",
|
||||
"model": "call-model",
|
||||
}));
|
||||
explicit = context().prepare_call(&explicit).unwrap();
|
||||
assert_eq!(explicit.arguments["model"], "call-model");
|
||||
|
||||
let default = task_call(serde_json::json!({
|
||||
"description": "task",
|
||||
"prompt": "inspect",
|
||||
"subagent_type": "generalPurpose",
|
||||
}));
|
||||
let default = context().prepare_call(&default).unwrap();
|
||||
assert_eq!(default.arguments["model"], "parent-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_renders_general_typed_and_custom_subagent_types_without_aliases() {
|
||||
let cases = [
|
||||
("generalPurpose", "unspecified"),
|
||||
("cursor-guide", "cursor-guide"),
|
||||
("MyReviewer", "MyReviewer"),
|
||||
];
|
||||
for (name, expected) in cases {
|
||||
let call = task_call(serde_json::json!({
|
||||
"description": "task",
|
||||
"prompt": "inspect",
|
||||
"subagent_type": name,
|
||||
}));
|
||||
let rendered = interaction::render_tool_call(&call, false).unwrap();
|
||||
let Some(pb::tool_call::Tool::TaskToolCall(tool)) = rendered.tool else {
|
||||
panic!("expected TaskToolCall")
|
||||
};
|
||||
let subagent = tool.args.unwrap().subagent_type.unwrap().r#type.unwrap();
|
||||
match (expected, subagent) {
|
||||
("unspecified", pb::subagent_type::Type::Unspecified(_)) => {}
|
||||
("cursor-guide", pb::subagent_type::Type::CursorGuide(_)) => {}
|
||||
(custom, pb::subagent_type::Type::Custom(value)) => {
|
||||
assert_eq!(value.name, custom)
|
||||
}
|
||||
_ => panic!("wrong subagent oneof for {name}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_task_model_is_left_for_the_model_visible_reminder() {
|
||||
let mut context = context();
|
||||
context
|
||||
.subagent_models
|
||||
.insert("security-review".into(), SubagentModel::Disabled);
|
||||
let call = task_call(serde_json::json!({
|
||||
"description": "review",
|
||||
"prompt": "inspect",
|
||||
"subagent_type": "security-review",
|
||||
}));
|
||||
assert!(context.task_disabled(&call));
|
||||
assert!(context
|
||||
.prepare_call(&call)
|
||||
.unwrap()
|
||||
.arguments
|
||||
.get("model")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn update_current_step_uses_a_one_based_turn_message_index() {
|
||||
let dispatcher = ToolDispatcher::new(CursorToolRuntime::default());
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "update-call".into(),
|
||||
model_call_id: "model-call".into(),
|
||||
name: "UpdateCurrentStep".into(),
|
||||
arguments_text: r#"{"current_step":"testing"}"#.into(),
|
||||
arguments: serde_json::json!({"current_step":"testing"}),
|
||||
};
|
||||
let completed = HashSet::new();
|
||||
let started = HashSet::new();
|
||||
let messages = vec![
|
||||
CanonicalMessage::text("old-runtime", Role::User, Origin::Runtime, "old turn"),
|
||||
CanonicalMessage {
|
||||
message_id: "old-assistant".into(),
|
||||
role: Role::Assistant,
|
||||
origin: Origin::Assistant,
|
||||
content: MessageContent::Assistant {
|
||||
text: "old response".into(),
|
||||
thinking: String::new(),
|
||||
tool_round_id: None,
|
||||
replay_state: None,
|
||||
tool_calls: Vec::new(),
|
||||
},
|
||||
runtime_event_id: None,
|
||||
},
|
||||
CanonicalMessage::text(
|
||||
"current-runtime",
|
||||
Role::User,
|
||||
Origin::Runtime,
|
||||
"current turn",
|
||||
),
|
||||
];
|
||||
let dispatched = dispatcher
|
||||
.start_batch(
|
||||
&[call],
|
||||
ToolBatchState {
|
||||
completed: &completed,
|
||||
started: &started,
|
||||
response_text: "",
|
||||
response_thinking: "",
|
||||
},
|
||||
&messages,
|
||||
&BTreeMap::new(),
|
||||
&context(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let completion = dispatched[0].completion.as_ref().unwrap();
|
||||
let Some(pb::tool_call::Tool::CommunicateUpdateToolCall(tool)) =
|
||||
completion.tool_call().tool.as_ref()
|
||||
else {
|
||||
panic!("expected CommunicateUpdateToolCall")
|
||||
};
|
||||
let pb::communicate_update_result::Result::Success(success) =
|
||||
tool.result.as_ref().unwrap().result.as_ref().unwrap()
|
||||
else {
|
||||
panic!("expected communicate update success")
|
||||
};
|
||||
assert_eq!(success.message_index, 1);
|
||||
assert_eq!(success.current_step, "testing");
|
||||
}
|
||||
|
||||
fn context() -> ExecContext {
|
||||
ExecContext {
|
||||
conversation_id: "child".into(),
|
||||
root_conversation_id: "root".into(),
|
||||
default_subagent_model: "parent-model".into(),
|
||||
subagent_models: HashMap::new(),
|
||||
allow_subagents: true,
|
||||
subagents_disabled: false,
|
||||
terminals_folder: "/tmp/terminals".into(),
|
||||
admin_command_denylist: Vec::new(),
|
||||
mcp_routes: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn task_call(arguments: serde_json::Value) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "task-call".into(),
|
||||
model_call_id: "model-call".into(),
|
||||
name: "Task".into(),
|
||||
arguments_text: arguments.to_string(),
|
||||
arguments,
|
||||
}
|
||||
}
|
||||
@@ -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,69 @@
|
||||
#![allow(dead_code)]
|
||||
|
||||
use std::{
|
||||
collections::VecDeque,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use cursor_server::{
|
||||
model::{ModelInvocation, ModelRequest},
|
||||
provider::{ModelEvent, Provider, ProviderStream},
|
||||
Error,
|
||||
};
|
||||
use futures_util::stream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
enum FakeResponse {
|
||||
Events(Vec<Result<ModelEvent, Error>>),
|
||||
Pending,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct FakeProvider {
|
||||
responses: Arc<Mutex<VecDeque<FakeResponse>>>,
|
||||
requests: Arc<Mutex<Vec<ModelRequest>>>,
|
||||
}
|
||||
|
||||
impl FakeProvider {
|
||||
pub fn push(&self, events: Vec<ModelEvent>) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Events(events.into_iter().map(Ok).collect()));
|
||||
}
|
||||
pub fn push_error(&self, error: Error) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Events(vec![Err(error)]));
|
||||
}
|
||||
pub fn push_pending(&self) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Pending);
|
||||
}
|
||||
pub fn requests(&self) -> Vec<ModelRequest> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Provider for FakeProvider {
|
||||
fn stream(
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
_cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
self.requests.lock().unwrap().push(invocation.request);
|
||||
let events = self
|
||||
.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop_front()
|
||||
.expect("fake response configured");
|
||||
match events {
|
||||
FakeResponse::Events(events) => Box::pin(stream::iter(events)),
|
||||
FakeResponse::Pending => Box::pin(stream::pending()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,358 @@
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::{connect, proto::agent::v1 as pb},
|
||||
cursor::{CursorCommand, CursorSessionRegistry},
|
||||
model::{
|
||||
ProjectedContent, ProviderEndpointInput, ProviderModelInput, ProviderType, Role, Usage,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
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 endpoint = store
|
||||
.create_provider(&ProviderEndpointInput {
|
||||
name: "Test".into(),
|
||||
provider_type: ProviderType::OpenAiChat,
|
||||
base_url: "https://example.com/v1".into(),
|
||||
api_key: None,
|
||||
custom_headers: serde_json::json!({}),
|
||||
extra_params: serde_json::json!({}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let configured_model = store
|
||||
.save_provider_model(
|
||||
endpoint.provider_id,
|
||||
&ProviderModelInput {
|
||||
model_id: "test-model".into(),
|
||||
display_name: "Test Model".into(),
|
||||
endpoint_type: ProviderType::OpenAiChat,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: None,
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ModelEvent::ThinkingStart,
|
||||
ModelEvent::ThinkingDelta("reason".into()),
|
||||
ModelEvent::ThinkingEnd,
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("hello".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(20_000),
|
||||
output_tokens: Some(8),
|
||||
total_tokens: Some(20_008),
|
||||
cache_read_tokens: Some(80),
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: Some(3),
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
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(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
);
|
||||
let handle = registry.get_or_create("request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(CursorCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run(
|
||||
"conversation",
|
||||
"hello",
|
||||
&configured_model.model_hash,
|
||||
)),
|
||||
})
|
||||
.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 token_deltas = Vec::new();
|
||||
let mut checkpoints = 0;
|
||||
let mut after_turn_checkpoints = Vec::new();
|
||||
let mut blobs = HashMap::<Vec<u8>, Vec<u8>>::new();
|
||||
let mut set_blob_ids = HashSet::new();
|
||||
let mut withheld_final_ack = 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 {
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = &kv.message {
|
||||
assert!(
|
||||
set_blob_ids.insert(set.blob_id.clone()),
|
||||
"an acknowledged content-addressed Blob must not be SET twice"
|
||||
);
|
||||
blobs.insert(set.blob_id.clone(), set.blob_data.clone());
|
||||
}
|
||||
if text == "hello" && !withheld_final_ack {
|
||||
withheld_final_ack = true;
|
||||
assert!(
|
||||
tokio::time::timeout(std::time::Duration::from_millis(100), output.recv())
|
||||
.await
|
||||
.is_err(),
|
||||
"TurnEnded/checkpoint must wait for the final Blob ACK"
|
||||
);
|
||||
}
|
||||
handle
|
||||
.command(CursorCommand::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(usage)) => {
|
||||
assert_eq!(usage.input_tokens, Some(20_000));
|
||||
assert_eq!(usage.output_tokens, Some(8));
|
||||
assert_eq!(usage.cache_read_tokens, Some(80));
|
||||
assert_eq!(usage.reasoning_tokens, Some(3));
|
||||
saw_turn_ended = true;
|
||||
}
|
||||
Some(pb::interaction_update::Message::TokenDelta(delta)) => {
|
||||
token_deltas.push(delta.tokens)
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
checkpoints += 1;
|
||||
if saw_turn_ended {
|
||||
after_turn_checkpoints.push(state);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert_eq!(text, "hello");
|
||||
assert_eq!(thinking, "reason");
|
||||
assert!(thinking_duration_ms.is_some_and(|duration| duration >= 1));
|
||||
assert!(saw_turn_ended);
|
||||
assert_eq!(token_deltas, vec![8]);
|
||||
assert!(withheld_final_ack);
|
||||
assert!(
|
||||
checkpoints >= 2,
|
||||
"final checkpoint is intentionally repeated"
|
||||
);
|
||||
assert_eq!(after_turn_checkpoints.len(), 3);
|
||||
assert_eq!(after_turn_checkpoints[0].pending_tool_calls.len(), 1);
|
||||
assert!(after_turn_checkpoints[1].pending_tool_calls.is_empty());
|
||||
assert_eq!(after_turn_checkpoints[1], after_turn_checkpoints[2]);
|
||||
let token_details = after_turn_checkpoints[1].token_details.as_ref().unwrap();
|
||||
assert_eq!(token_details.used_tokens, 20_008);
|
||||
assert_eq!(token_details.max_tokens, 256_000);
|
||||
let breakdown = token_details.breakdown.as_ref().unwrap();
|
||||
assert_eq!(breakdown.total_used_tokens, 20_008);
|
||||
assert_eq!(breakdown.max_tokens, 256_000);
|
||||
assert_eq!(breakdown.categories.len(), 9);
|
||||
assert_eq!(
|
||||
breakdown
|
||||
.categories
|
||||
.iter()
|
||||
.map(|category| category.estimated_tokens)
|
||||
.sum::<u32>(),
|
||||
20_009
|
||||
);
|
||||
assert_eq!(
|
||||
after_turn_checkpoints[0].turns, after_turn_checkpoints[1].turns,
|
||||
"staged and settled checkpoints reuse the same frozen Turn"
|
||||
);
|
||||
assert_eq!(
|
||||
after_turn_checkpoints[0].root_prompt_messages_json.len() + 1,
|
||||
after_turn_checkpoints[1].root_prompt_messages_json.len()
|
||||
);
|
||||
let pending: serde_json::Value =
|
||||
serde_json::from_str(&after_turn_checkpoints[0].pending_tool_calls[0]).unwrap();
|
||||
assert_eq!(pending["role"], "assistant");
|
||||
assert!(pending["content"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.any(|part| part["type"] == "text" && part["text"] == "hello"));
|
||||
let final_root = after_turn_checkpoints[1]
|
||||
.root_prompt_messages_json
|
||||
.last()
|
||||
.unwrap();
|
||||
let stable: serde_json::Value = serde_json::from_slice(blobs.get(final_root).unwrap()).unwrap();
|
||||
assert_eq!(stable["role"], "assistant");
|
||||
assert!(
|
||||
stable.get("origin").is_none(),
|
||||
"wire root is not CanonicalMessage JSON"
|
||||
);
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
blobs
|
||||
.get(after_turn_checkpoints[0].turns.last().unwrap())
|
||||
.unwrap()
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent turn")
|
||||
};
|
||||
assert_eq!(turn.steps.len(), 2, "thinking and text are frozen once");
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert!(requests[0]
|
||||
.prompt
|
||||
.instructions
|
||||
.contains("powered by Test Model"));
|
||||
let projected = &requests[0].history;
|
||||
assert_eq!(projected[0].role, Role::User);
|
||||
let ProjectedContent::Parts(runtime) = &projected[0].content else {
|
||||
panic!("runtime context must be text")
|
||||
};
|
||||
assert!(matches!(
|
||||
runtime.as_slice(),
|
||||
[cursor_server::model::ContentPart::Text { text }]
|
||||
if text.contains("<user_info>")
|
||||
&& text.contains("<user_query>\nhello\n</user_query>")
|
||||
));
|
||||
assert_eq!(
|
||||
projected.len(),
|
||||
1,
|
||||
"the raw UserMessage is not projected twice"
|
||||
);
|
||||
|
||||
let messages = store
|
||||
.load_current_messages(&cursor_server::model::ConversationId::new("conversation"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(messages[0].message_id, "runtime:run-request:request");
|
||||
assert_eq!(messages[0].role, Role::User);
|
||||
assert_eq!(messages.len(), 2, "runtime user plus final assistant");
|
||||
let stored_runs: Vec<String> = sqlx::query_scalar("SELECT run_id FROM runs ORDER BY run_id")
|
||||
.fetch_all(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
stored_runs,
|
||||
vec!["request"],
|
||||
"the concrete request_id, not Cursor's reusable wire run_id, owns the execution"
|
||||
);
|
||||
}
|
||||
|
||||
fn client_run(conversation_id: &str, text: &str, model_id: &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 {
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: model_id.into(),
|
||||
parameters: vec![pb::requested_model::ModelParameterValue {
|
||||
id: "context".into(),
|
||||
value: "256k".into(),
|
||||
}],
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
request_context: Some(pb::RequestContext {
|
||||
env: Some(pb::RequestContextEnv {
|
||||
os_version: "darwin".into(),
|
||||
workspace_paths: vec!["/workspace".into()],
|
||||
shell: "zsh".into(),
|
||||
terminals_folder: "/terminals".into(),
|
||||
agent_transcripts_folder: "/transcripts".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
git_repos: vec![pb::GitRepoInfo {
|
||||
path: "/workspace".into(),
|
||||
status: "M src/main.rs".into(),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some(conversation_id.into()),
|
||||
conversation_state: None,
|
||||
run_id: Some("reusable-wire-run-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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,117 @@
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use cursor_server::{
|
||||
model::{
|
||||
ConversationId, MessageContent, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId,
|
||||
RunKind, ToolCall, ToolResult, ToolRoundAssistant, ToolRoundId,
|
||||
},
|
||||
store::ToolRoundStatus,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn results_commit_adjacent_pairs_in_arrival_order() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let conversation_id = ConversationId::new("conversation");
|
||||
let root = store.ensure_conversation(&conversation_id).await.unwrap();
|
||||
let run = PreparedRun {
|
||||
run_id: RunId::new("run"),
|
||||
conversation_id: conversation_id.clone(),
|
||||
kind: RunKind::Root,
|
||||
model: ModelSpec::new("model"),
|
||||
prompt: PromptSpec {
|
||||
instructions: String::new(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
initial_messages: Vec::new(),
|
||||
action: RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
},
|
||||
base_revision_id: root,
|
||||
};
|
||||
store.claim_run(&run).await.unwrap();
|
||||
let round_id = ToolRoundId::new("round");
|
||||
let calls = [call(0, "A"), call(1, "B"), call(2, "C")];
|
||||
store
|
||||
.create_tool_round(
|
||||
&round_id,
|
||||
&run.run_id,
|
||||
root,
|
||||
&ToolRoundAssistant {
|
||||
text: "answer prefix".into(),
|
||||
thinking: "reasoning".into(),
|
||||
model_call_id: "model-call".into(),
|
||||
replay_state: None,
|
||||
},
|
||||
&calls,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let b = store
|
||||
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("B"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(b.completion_seq, 0);
|
||||
assert!(!b.settled);
|
||||
let a = store
|
||||
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("A"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(a.completion_seq, 1);
|
||||
let c = store
|
||||
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("C"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(c.settled);
|
||||
|
||||
let messages = store.load_revision_messages(c.revision_id).await.unwrap();
|
||||
assert_eq!(messages.len(), 6);
|
||||
let ids = messages
|
||||
.chunks_exact(2)
|
||||
.map(|pair| match (&pair[0].content, &pair[1].content) {
|
||||
(MessageContent::Assistant { tool_calls, .. }, MessageContent::ToolResult(result)) => {
|
||||
assert_eq!(tool_calls[0].call_id, result.call_id);
|
||||
result.call_id.clone()
|
||||
}
|
||||
_ => panic!("tool result must be adjacent to its assistant call"),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(ids, ["B", "A", "C"]);
|
||||
let MessageContent::Assistant { text, thinking, .. } = &messages[0].content else {
|
||||
unreachable!()
|
||||
};
|
||||
assert_eq!(text, "answer prefix");
|
||||
assert_eq!(thinking, "reasoning");
|
||||
for message in [&messages[2], &messages[4]] {
|
||||
let MessageContent::Assistant { text, thinking, .. } = &message.content else {
|
||||
unreachable!()
|
||||
};
|
||||
assert!(text.is_empty());
|
||||
assert!(thinking.is_empty());
|
||||
}
|
||||
let snapshot = store.tool_round(&round_id).await.unwrap().unwrap();
|
||||
assert_eq!(snapshot.status, ToolRoundStatus::Settled);
|
||||
assert_eq!(snapshot.completed_call_ids.len(), 3);
|
||||
}
|
||||
|
||||
fn call(index: usize, call_id: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
model_call_id: "model-call".into(),
|
||||
name: "Read".into(),
|
||||
arguments_text: r#"{"path":"/tmp/a"}"#.into(),
|
||||
arguments: serde_json::json!({"path":"/tmp/a"}),
|
||||
}
|
||||
}
|
||||
|
||||
fn result(call_id: &str) -> ToolResult {
|
||||
ToolResult {
|
||||
call_id: call_id.into(),
|
||||
content: format!("result-{call_id}"),
|
||||
is_error: false,
|
||||
image: None,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
use axum::{http::StatusCode, response::IntoResponse, routing::get, Router};
|
||||
use cursor_server::web::{HtmlEngine, JsonEngine, WebSearch};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
const RESULT_SELECTOR: &str = ".result";
|
||||
const TITLE_SELECTOR: &str = ".title";
|
||||
const LINK_SELECTOR: &str = "a.title";
|
||||
const SNIPPET_SELECTOR: &str = ".snippet";
|
||||
|
||||
#[test]
|
||||
fn built_in_catalog_covers_the_reference_search_engines() {
|
||||
let search = WebSearch::built_in();
|
||||
let ids = search.engine_ids();
|
||||
for expected in [
|
||||
"google",
|
||||
"bing",
|
||||
"brave",
|
||||
"duckduckgo",
|
||||
"startpage",
|
||||
"yahoo",
|
||||
"mojeek",
|
||||
"qwant",
|
||||
"ecosia",
|
||||
"yandex",
|
||||
"baidu",
|
||||
"sogou",
|
||||
"so360",
|
||||
"naver",
|
||||
"seznam",
|
||||
"wikipedia",
|
||||
"github",
|
||||
"stackoverflow",
|
||||
"crates_io",
|
||||
"npm",
|
||||
"pypi",
|
||||
"arxiv",
|
||||
"crossref",
|
||||
] {
|
||||
assert!(ids.contains(&expected), "missing engine {expected}");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn federated_search_deduplicates_and_rrf_ranks_results() {
|
||||
let base = spawn_search_fixture().await;
|
||||
let search = WebSearch::with_engines(vec![
|
||||
engine("first", format!("{base}/first?q={{query}}")),
|
||||
engine("second", format!("{base}/second?q={{query}}")),
|
||||
]);
|
||||
|
||||
let results = search.search("rust agent").await.unwrap();
|
||||
|
||||
assert_eq!(results.len(), 3);
|
||||
assert_eq!(results[0].url, "https://example.com/shared");
|
||||
assert_eq!(results[0].engines, vec!["first", "second"]);
|
||||
assert!(results.iter().any(|result| result.title == "Alpha"));
|
||||
assert!(results.iter().any(|result| result.title == "Beta"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn one_failed_engine_does_not_discard_other_engine_results() {
|
||||
let base = spawn_search_fixture().await;
|
||||
let search = WebSearch::with_engines(vec![
|
||||
engine("failed", format!("{base}/failed?q={{query}}")),
|
||||
engine("first", format!("{base}/first?q={{query}}")),
|
||||
]);
|
||||
|
||||
let results = search.search("rust").await.unwrap();
|
||||
|
||||
assert_eq!(results.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn all_failed_engines_return_a_search_error() {
|
||||
let base = spawn_search_fixture().await;
|
||||
let search =
|
||||
WebSearch::with_engines(vec![engine("failed", format!("{base}/failed?q={{query}}"))]);
|
||||
|
||||
let error = search.search("rust").await.unwrap_err();
|
||||
|
||||
assert!(error.to_string().contains("failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn json_engines_use_declared_result_fields() {
|
||||
let base = spawn_search_fixture().await;
|
||||
let search = WebSearch::with_engines(vec![JsonEngine::new(
|
||||
"json",
|
||||
format!("{base}/json?q={{query}}"),
|
||||
"/items",
|
||||
"/name",
|
||||
"/url",
|
||||
"/description",
|
||||
None,
|
||||
)]);
|
||||
|
||||
let results = search.search("rust").await.unwrap();
|
||||
|
||||
assert_eq!(results.len(), 1);
|
||||
assert_eq!(results[0].title, "Structured result");
|
||||
assert_eq!(results[0].url, "https://example.com/structured");
|
||||
assert_eq!(results[0].chunk, "Parsed from JSON");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "live public search smoke test"]
|
||||
async fn built_in_search_returns_live_results() {
|
||||
let _ = tracing_subscriber::fmt()
|
||||
.with_env_filter("cursor_server::web=debug")
|
||||
.try_init();
|
||||
let results = WebSearch::built_in()
|
||||
.search("Rust programming language")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(!results.is_empty());
|
||||
for result in results {
|
||||
println!("{}\t{}\t{:?}", result.title, result.url, result.engines);
|
||||
}
|
||||
}
|
||||
|
||||
fn engine(id: &'static str, url: String) -> HtmlEngine {
|
||||
HtmlEngine::new(
|
||||
id,
|
||||
url,
|
||||
RESULT_SELECTOR,
|
||||
TITLE_SELECTOR,
|
||||
LINK_SELECTOR,
|
||||
SNIPPET_SELECTOR,
|
||||
)
|
||||
}
|
||||
|
||||
async fn spawn_search_fixture() -> String {
|
||||
async fn first() -> impl IntoResponse {
|
||||
r#"
|
||||
<article class="result">
|
||||
<a class="title" href="https://example.com/shared?utm_source=one">Shared</a>
|
||||
<p class="snippet">Shared from first.</p>
|
||||
</article>
|
||||
<article class="result">
|
||||
<a class="title" href="https://alpha.example/path">Alpha</a>
|
||||
<p class="snippet">Alpha result.</p>
|
||||
</article>
|
||||
"#
|
||||
}
|
||||
async fn second() -> impl IntoResponse {
|
||||
r#"
|
||||
<article class="result">
|
||||
<a class="title" href="https://example.com/shared#section">Shared result</a>
|
||||
<p class="snippet">Shared from second.</p>
|
||||
</article>
|
||||
<article class="result">
|
||||
<a class="title" href="https://beta.example/">Beta</a>
|
||||
<p class="snippet">Beta result.</p>
|
||||
</article>
|
||||
"#
|
||||
}
|
||||
let app = Router::new()
|
||||
.route("/first", get(first))
|
||||
.route("/second", get(second))
|
||||
.route(
|
||||
"/failed",
|
||||
get(|| async { (StatusCode::TOO_MANY_REQUESTS, "limited") }),
|
||||
)
|
||||
.route(
|
||||
"/json",
|
||||
get(|| async {
|
||||
axum::Json(serde_json::json!({
|
||||
"items": [{
|
||||
"name": "Structured result",
|
||||
"url": "https://example.com/structured",
|
||||
"description": "Parsed from JSON"
|
||||
}]
|
||||
}))
|
||||
}),
|
||||
);
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
|
||||
format!("http://{address}")
|
||||
}
|
||||
Reference in New Issue
Block a user