mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
feat: integrate app version into control service and update ads endpoint
- Added `app_version` field to `ControlService` and updated its initialization to include the app version. - Modified the `ADS_ENDPOINT` to point to a local server for development purposes. - Refactored conversation command and output handling to utilize a new `RunFinish` enum for better state management. - Enhanced the conversation runtime to handle queued user messages after a turn has ended, ensuring smooth transitions between turns. - Added tests to validate the new behavior of queued messages and transport handling.
This commit is contained in:
@@ -66,6 +66,7 @@ impl App {
|
||||
plugin_runtime,
|
||||
plugins,
|
||||
clients.clone(),
|
||||
config.app_version.clone(),
|
||||
)?;
|
||||
let harness = control.cursor_harness().clone();
|
||||
let mut router = api::router(registry.clone(), clients)?;
|
||||
|
||||
@@ -16,7 +16,8 @@ use super::ControlService;
|
||||
// 此广告拉取不涉及用户隐私,用户id随机产生
|
||||
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
|
||||
|
||||
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
|
||||
// pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
|
||||
pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu";
|
||||
pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
|
||||
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
||||
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
||||
|
||||
@@ -42,6 +42,7 @@ pub struct ControlService {
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
clients: crate::network::NetworkClients,
|
||||
app_version: String,
|
||||
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
||||
}
|
||||
|
||||
@@ -153,6 +154,7 @@ impl ControlService {
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
clients: crate::network::NetworkClients,
|
||||
app_version: String,
|
||||
) -> Result<Self> {
|
||||
Ok(Self {
|
||||
cursor_harness: CursorHarness::new(store.clone())?,
|
||||
@@ -161,6 +163,7 @@ impl ControlService {
|
||||
plugin_runtime,
|
||||
plugins,
|
||||
clients,
|
||||
app_version,
|
||||
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
||||
})
|
||||
}
|
||||
@@ -265,7 +268,7 @@ impl ControlService {
|
||||
.get(ADS_ENDPOINT)
|
||||
.header(DEVICE_ID_HEADER, installation_id)
|
||||
.header(OS_HEADER, std::env::consts::OS)
|
||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
||||
.header(APP_VERSION_HEADER, &self.app_version)
|
||||
.header(LANGUAGE_HEADER, language)
|
||||
.timeout(std::time::Duration::from_secs(60));
|
||||
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
|
||||
@@ -299,7 +302,7 @@ impl ControlService {
|
||||
.post(endpoint)
|
||||
.header(DEVICE_ID_HEADER, installation_id)
|
||||
.header(OS_HEADER, std::env::consts::OS)
|
||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
||||
.header(APP_VERSION_HEADER, &self.app_version)
|
||||
.json(input)
|
||||
.timeout(std::time::Duration::from_secs(5))
|
||||
.send()
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
|
||||
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum RunFinish {
|
||||
TurnCompleted,
|
||||
Transport(TransportFinish),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum TransportFinish {
|
||||
Success,
|
||||
@@ -17,7 +23,7 @@ pub enum TransportCommand {
|
||||
},
|
||||
RunFinished {
|
||||
generation: u64,
|
||||
finish: TransportFinish,
|
||||
finish: RunFinish,
|
||||
},
|
||||
Disconnect,
|
||||
}
|
||||
|
||||
@@ -34,7 +34,7 @@ use crate::{
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, TransportFinish};
|
||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
|
||||
use crate::cursor::transport::TransportHandle;
|
||||
|
||||
pub struct ConversationOutput {
|
||||
@@ -110,7 +110,7 @@ impl ConversationOutput {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(mut self) -> Result<TransportFinish> {
|
||||
pub async fn run(mut self) -> Result<RunFinish> {
|
||||
let result = self.run_inner().await;
|
||||
if let Err(error) = &result {
|
||||
if !self.superseded.is_cancelled() {
|
||||
@@ -143,7 +143,7 @@ impl ConversationOutput {
|
||||
result
|
||||
}
|
||||
|
||||
async fn run_inner(&mut self) -> Result<TransportFinish> {
|
||||
async fn run_inner(&mut self) -> Result<RunFinish> {
|
||||
if self.context.compacting {
|
||||
self.handle.emit(&events::summary_started())?;
|
||||
}
|
||||
@@ -175,7 +175,7 @@ impl ConversationOutput {
|
||||
if self.superseded.is_cancelled() {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(TransportFinish::Cancelled);
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
let input = if let Ok(action) = self.runtime_actions.try_recv() {
|
||||
Input::RuntimeAction(Some(Box::new(action)))
|
||||
@@ -187,7 +187,7 @@ impl ConversationOutput {
|
||||
_ = self.superseded.cancelled() => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(TransportFinish::Cancelled);
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
||||
event = self.core.events.recv() => Input::Event(event),
|
||||
@@ -719,7 +719,7 @@ impl ConversationOutput {
|
||||
if self.superseded.is_cancelled() {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(TransportFinish::Cancelled);
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
return match outcome {
|
||||
RunOutcome::Completed => {
|
||||
@@ -735,7 +735,7 @@ impl ConversationOutput {
|
||||
for _ in 0..3 {
|
||||
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
||||
}
|
||||
return Ok(TransportFinish::Success);
|
||||
return Ok(RunFinish::TurnCompleted);
|
||||
}
|
||||
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
||||
Error::Protocol("Completed without final state".into())
|
||||
@@ -751,17 +751,19 @@ impl ConversationOutput {
|
||||
ttft_breakdown: None,
|
||||
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
||||
})?;
|
||||
Ok(TransportFinish::Success)
|
||||
Ok(RunFinish::TurnCompleted)
|
||||
}
|
||||
RunOutcome::Cancelled => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
Ok(TransportFinish::Cancelled)
|
||||
Ok(RunFinish::Transport(TransportFinish::Cancelled))
|
||||
}
|
||||
RunOutcome::Failed(failure) => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
Ok(TransportFinish::Failed(cursor_error(failure)))
|
||||
Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error(
|
||||
failure,
|
||||
))))
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -18,19 +18,22 @@ use crate::{
|
||||
},
|
||||
transport::{OrderedInbox, TransportHandle},
|
||||
},
|
||||
run::{CommandResult, RunEngine, RunHandle},
|
||||
run::{CommandResult, RunEngine, RunHandle, RunPhase},
|
||||
};
|
||||
|
||||
use super::{
|
||||
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
||||
ConversationRegistry, MessageDelivery, TransportCommand, TransportFinish,
|
||||
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
|
||||
};
|
||||
|
||||
pub struct ConversationRuntime;
|
||||
|
||||
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
|
||||
|
||||
#[derive(Clone)]
|
||||
struct RunGeneration {
|
||||
id: u64,
|
||||
request: pb::AgentRunRequest,
|
||||
superseded: CancellationToken,
|
||||
finished: CancellationToken,
|
||||
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
||||
@@ -78,6 +81,7 @@ impl ConversationRuntime {
|
||||
let mut next_generation = 1_u64;
|
||||
let mut pending_finish = None::<(u64, TransportFinish)>;
|
||||
let mut draining = false;
|
||||
let mut waiting_for_action = false;
|
||||
loop {
|
||||
let command = if draining {
|
||||
if !handle.admissions_drained() {
|
||||
@@ -102,6 +106,28 @@ impl ConversationRuntime {
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if waiting_for_action {
|
||||
tokio::select! {
|
||||
command = receiver.recv() => match command {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
}
|
||||
},
|
||||
_ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
};
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation.id, TransportFinish::Success));
|
||||
draining = true;
|
||||
waiting_for_action = false;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
match receiver.recv().await {
|
||||
Some(command) => command,
|
||||
@@ -130,7 +156,19 @@ impl ConversationRuntime {
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
}
|
||||
let turn_completed = waiting_for_action
|
||||
|| current.as_ref().is_some_and(|generation| {
|
||||
generation
|
||||
.run
|
||||
.lock()
|
||||
.as_ref()
|
||||
.is_none_or(|run| run.phase() != RunPhase::Running)
|
||||
});
|
||||
if turn_completed {
|
||||
super::finish_success(&handle);
|
||||
} else {
|
||||
super::finish_cancelled(&handle).ok();
|
||||
}
|
||||
break;
|
||||
}
|
||||
TransportCommand::RunFinished { generation, finish } => {
|
||||
@@ -140,10 +178,20 @@ impl ConversationRuntime {
|
||||
{
|
||||
continue;
|
||||
}
|
||||
match finish {
|
||||
RunFinish::TurnCompleted => {
|
||||
pending_finish = None;
|
||||
draining = false;
|
||||
waiting_for_action = true;
|
||||
}
|
||||
RunFinish::Transport(finish) => {
|
||||
waiting_for_action = false;
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation, finish));
|
||||
draining = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
TransportCommand::Append { seqno, message } => {
|
||||
for (_seqno, message) in inbox.push(seqno, *message) {
|
||||
{
|
||||
@@ -151,6 +199,7 @@ impl ConversationRuntime {
|
||||
Some(pb::agent_client_message::Message::RunRequest(
|
||||
request,
|
||||
)) => {
|
||||
waiting_for_action = false;
|
||||
if draining {
|
||||
handle.reopen();
|
||||
draining = false;
|
||||
@@ -171,57 +220,18 @@ impl ConversationRuntime {
|
||||
return;
|
||||
}
|
||||
}
|
||||
let previous_finished =
|
||||
if let Some(previous) = current.take() {
|
||||
previous.superseded.cancel();
|
||||
if let Some(run) = previous.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
for id in previous
|
||||
.tool_runtime
|
||||
.interrupt_for_run_replacement()
|
||||
.await
|
||||
{
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
Some(previous.finished.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (results, result_receiver) = tool_result_channel();
|
||||
let (runtime_actions, runtime_action_receiver) =
|
||||
mpsc::unbounded_channel::<compile::RuntimeAction>();
|
||||
let tool_runtime = tool_runtime_factory.next_run();
|
||||
let tools = ToolDispatcher::with_results(
|
||||
tool_runtime.clone(),
|
||||
results.clone(),
|
||||
dependencies.store.clone(),
|
||||
dependencies.web_cache.clone(),
|
||||
);
|
||||
let generation = RunGeneration {
|
||||
id: next_generation,
|
||||
superseded: CancellationToken::new(),
|
||||
finished: CancellationToken::new(),
|
||||
run: Arc::new(parking_lot::Mutex::new(None)),
|
||||
results,
|
||||
runtime_actions,
|
||||
tool_runtime,
|
||||
tools,
|
||||
};
|
||||
next_generation = next_generation.saturating_add(1);
|
||||
current = Some(generation.clone());
|
||||
spawn_run_request(
|
||||
registry.clone(),
|
||||
handle.clone(),
|
||||
start_generation(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
dependencies.clone(),
|
||||
blob_sync.clone(),
|
||||
context_sync.clone(),
|
||||
generation,
|
||||
previous_finished,
|
||||
result_receiver,
|
||||
runtime_action_receiver,
|
||||
);
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
message,
|
||||
@@ -382,27 +392,48 @@ impl ConversationRuntime {
|
||||
// return an explicit Protocol Error rather than falling through silently.
|
||||
Some(
|
||||
pb::agent_client_message::Message::ConversationAction(
|
||||
action,
|
||||
conversation_action,
|
||||
),
|
||||
) => match action.action {
|
||||
) => match conversation_action.action.clone() {
|
||||
Some(
|
||||
pb::conversation_action::Action::UserMessageAction(
|
||||
action,
|
||||
),
|
||||
) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
let delivered_to_active_run =
|
||||
current.as_ref().is_some_and(|generation| {
|
||||
generation.run.lock().as_ref().is_some_and(
|
||||
|run| run.phase() == RunPhase::Running,
|
||||
) && generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(
|
||||
action.clone(),
|
||||
))
|
||||
.is_ok()
|
||||
});
|
||||
if delivered_to_active_run {
|
||||
continue;
|
||||
}
|
||||
let Some(previous) = current.as_ref() else {
|
||||
continue;
|
||||
};
|
||||
if generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(action))
|
||||
.is_err()
|
||||
{
|
||||
generation.results.send_error(crate::Error::Protocol(
|
||||
"UserMessageAction arrived without an active Run"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
let mut request = previous.request.clone();
|
||||
request.action = Some(conversation_action);
|
||||
request.conversation_state = None;
|
||||
request.pre_fetched_blobs.clear();
|
||||
waiting_for_action = false;
|
||||
start_generation(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||
if let Some(generation) = current.as_ref() {
|
||||
@@ -479,6 +510,67 @@ impl ConversationRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn start_generation(
|
||||
registry: &ConversationRegistry,
|
||||
handle: &TransportHandle,
|
||||
dependencies: &ConversationDependencies,
|
||||
blob_sync: &BlobSynchronizer,
|
||||
context_sync: &RequestContextSynchronizer,
|
||||
tool_runtime_factory: &CursorToolRuntime,
|
||||
current: &mut Option<RunGeneration>,
|
||||
next_generation: &mut u64,
|
||||
request: pb::AgentRunRequest,
|
||||
) {
|
||||
let previous_finished = if let Some(previous) = current.take() {
|
||||
previous.superseded.cancel();
|
||||
if let Some(run) = previous.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
for id in previous.tool_runtime.interrupt_for_run_replacement().await {
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
Some(previous.finished.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (results, result_receiver) = tool_result_channel();
|
||||
let (runtime_actions, runtime_action_receiver) =
|
||||
mpsc::unbounded_channel::<compile::RuntimeAction>();
|
||||
let tool_runtime = tool_runtime_factory.next_run();
|
||||
let tools = ToolDispatcher::with_results(
|
||||
tool_runtime.clone(),
|
||||
results.clone(),
|
||||
dependencies.store.clone(),
|
||||
dependencies.web_cache.clone(),
|
||||
);
|
||||
let generation = RunGeneration {
|
||||
id: *next_generation,
|
||||
request: request.clone(),
|
||||
superseded: CancellationToken::new(),
|
||||
finished: CancellationToken::new(),
|
||||
run: Arc::new(parking_lot::Mutex::new(None)),
|
||||
results,
|
||||
runtime_actions,
|
||||
tool_runtime,
|
||||
tools,
|
||||
};
|
||||
*next_generation = next_generation.saturating_add(1);
|
||||
*current = Some(generation.clone());
|
||||
spawn_run_request(
|
||||
registry.clone(),
|
||||
handle.clone(),
|
||||
request,
|
||||
dependencies.clone(),
|
||||
blob_sync.clone(),
|
||||
context_sync.clone(),
|
||||
generation,
|
||||
previous_finished,
|
||||
result_receiver,
|
||||
runtime_action_receiver,
|
||||
);
|
||||
}
|
||||
|
||||
fn finish_pending(
|
||||
handle: &TransportHandle,
|
||||
current: &Option<RunGeneration>,
|
||||
@@ -569,7 +661,7 @@ fn spawn_run_request(
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: TransportFinish::Failed(error),
|
||||
finish: RunFinish::Transport(TransportFinish::Failed(error)),
|
||||
})
|
||||
.await;
|
||||
return;
|
||||
@@ -606,7 +698,7 @@ fn spawn_run_request(
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: TransportFinish::Success,
|
||||
finish: RunFinish::Transport(TransportFinish::Success),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
@@ -635,7 +727,7 @@ fn spawn_run_request(
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: TransportFinish::Success,
|
||||
finish: RunFinish::Transport(TransportFinish::Success),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
@@ -700,14 +792,14 @@ fn spawn_run_request(
|
||||
Ok(finish) => finish,
|
||||
Err(error) => {
|
||||
if generation.superseded.is_cancelled() {
|
||||
TransportFinish::Cancelled
|
||||
RunFinish::Transport(TransportFinish::Cancelled)
|
||||
} else {
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"Cursor session failed"
|
||||
);
|
||||
TransportFinish::Failed(error)
|
||||
RunFinish::Transport(TransportFinish::Failed(error))
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -422,6 +422,85 @@ async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() {
|
||||
assert_eq!(output.recv().await, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn queued_user_message_after_turn_ended_starts_the_next_turn() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(text_response("first turn"));
|
||||
provider.push(text_response("queued turn"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("queued-after-turn").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run_for(
|
||||
"queued-after-turn",
|
||||
"queued-after-turn-conversation",
|
||||
)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
wait_for_turn_ended(&handle, &mut output, &mut append_seqno).await;
|
||||
assert_transport_remains_open(&handle, &mut output, &mut append_seqno).await;
|
||||
|
||||
cursor_server::api::cursor::bidi::append(
|
||||
®istry,
|
||||
cursor_server::api::cursor::bidi::DecodedAppend {
|
||||
request_id: "queued-after-turn".into(),
|
||||
seqno: append_seqno,
|
||||
message: runtime_user_message(),
|
||||
},
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
|
||||
let mut text = String::new();
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("queued turn closed without 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 {
|
||||
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
|
||||
text.push_str(&delta.text);
|
||||
}
|
||||
}
|
||||
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
|
||||
}
|
||||
|
||||
assert!(text.contains("queued turn"));
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
assert_eq!(
|
||||
&requests[1].history[..requests[0].history.len()],
|
||||
requests[0].history.as_slice(),
|
||||
"queued continuation must preserve the first provider request as a prefix"
|
||||
);
|
||||
let history = serde_json::to_string(&requests[1].history).unwrap();
|
||||
assert!(history.contains("queued follow-up"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_user_message_action_interrupts_and_continues_with_new_message() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
@@ -1732,6 +1811,55 @@ async fn run_to_end(
|
||||
}
|
||||
}
|
||||
|
||||
async fn wait_for_turn_ended(
|
||||
handle: &cursor_server::cursor::TransportHandle,
|
||||
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
|
||||
append_seqno: &mut i64,
|
||||
) {
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before turnEnded");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
let turn_ended = matches!(
|
||||
server.message,
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(
|
||||
pb::InteractionUpdate {
|
||||
message: Some(pb::interaction_update::Message::TurnEnded(_)),
|
||||
}
|
||||
))
|
||||
);
|
||||
acknowledge_kv(handle, append_seqno, &frame).await;
|
||||
if turn_ended {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn assert_transport_remains_open(
|
||||
handle: &cursor_server::cursor::TransportHandle,
|
||||
output: &mut tokio::sync::mpsc::UnboundedReceiver<Bytes>,
|
||||
append_seqno: &mut i64,
|
||||
) {
|
||||
let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(100);
|
||||
loop {
|
||||
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
|
||||
let Ok(Some(frame)) = tokio::time::timeout(remaining, output.recv()).await else {
|
||||
return;
|
||||
};
|
||||
let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
assert_eq!(
|
||||
flags & connect::END_STREAM_FLAG,
|
||||
0,
|
||||
"turnEnded closed the transport before the queued action arrived"
|
||||
);
|
||||
acknowledge_kv(handle, append_seqno, &frame).await;
|
||||
}
|
||||
}
|
||||
|
||||
async fn acknowledge_kv(
|
||||
handle: &cursor_server::cursor::TransportHandle,
|
||||
append_seqno: &mut i64,
|
||||
|
||||
Reference in New Issue
Block a user