mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
Merge remote-tracking branch 'origin/main'
This commit is contained in:
@@ -158,6 +158,7 @@ impl ConversationOutput {
|
||||
let mut streams = BTreeMap::<usize, ToolCallStream>::new();
|
||||
let mut completions = HashMap::<String, ToolCompletion>::new();
|
||||
let mut completed = HashSet::<String>::new();
|
||||
let mut completed_round = None::<ToolRoundId>;
|
||||
let mut response_text = String::new();
|
||||
let mut response_thinking = String::new();
|
||||
let mut active_round = None::<ToolRoundId>;
|
||||
@@ -425,6 +426,16 @@ impl ConversationOutput {
|
||||
round_id,
|
||||
calls: round_calls,
|
||||
} => {
|
||||
// `completed` exists so that replaying ExecuteToolRound for a round
|
||||
// does not dispatch a call this round already committed. Tool call ids
|
||||
// are only unique *within* a round -- the schema says as much with
|
||||
// `UNIQUE (round_id, call_id)` -- so an id retained from an earlier
|
||||
// round would make start_batch skip a fresh call, and tool_round wait
|
||||
// forever for a result nothing will ever produce.
|
||||
if completed_round.as_ref() != Some(&round_id) {
|
||||
completed.clear();
|
||||
completed_round = Some(round_id.clone());
|
||||
}
|
||||
active_round = Some(round_id.clone());
|
||||
active_tool_calls = round_calls
|
||||
.iter()
|
||||
|
||||
@@ -1253,6 +1253,151 @@ fn client_run_for_model(
|
||||
}
|
||||
}
|
||||
|
||||
/// A provider that reuses a tool call id across two rounds of the same run
|
||||
/// must not wedge the run. `ToolDispatcher::start_batch` skips any call whose
|
||||
/// id is already in `ToolBatchState::completed`, and that set is never cleared
|
||||
/// for the lifetime of the run, so the skipped call produced no completion and
|
||||
/// `tool_round::execute` waited forever for a result that could never arrive.
|
||||
#[tokio::test]
|
||||
async fn duplicate_tool_call_id_across_rounds_does_not_wedge_the_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
for _ in 0..2 {
|
||||
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),
|
||||
]);
|
||||
}
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "ignored".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("done".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 = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("tool-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut seqno = 1;
|
||||
let mut execs = 0;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(10), output.recv())
|
||||
.await
|
||||
.expect("run must not hang on a reused tool call id")
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
execs += 1;
|
||||
let exec_id = exec.id;
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: exec_id,
|
||||
exec_id: String::new(),
|
||||
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||
pb::ReadResult {
|
||||
result: Some(pb::read_result::Result::Success(
|
||||
pb::ReadSuccess {
|
||||
path: "/tmp/a".into(),
|
||||
total_lines: 1,
|
||||
file_size: 1,
|
||||
output: Some(
|
||||
pb::read_success::Output::Content(
|
||||
"x".into(),
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(
|
||||
pb::agent_client_message::Message::ExecClientControlMessage(
|
||||
pb::ExecClientControlMessage {
|
||||
message: Some(
|
||||
pb::exec_client_control_message::Message::StreamClose(
|
||||
pb::ExecClientStreamClose { id: exec_id },
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
// Both rounds have to reach the client, and the run has to get far enough
|
||||
// to ask the provider a third time and finish.
|
||||
assert_eq!(execs, 2);
|
||||
assert_eq!(provider.requests().len(), 3);
|
||||
}
|
||||
|
||||
fn client_run() -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: "read it".into(),
|
||||
|
||||
Reference in New Issue
Block a user