mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-09 00:16:23 +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,
|
plugin_runtime,
|
||||||
plugins,
|
plugins,
|
||||||
clients.clone(),
|
clients.clone(),
|
||||||
|
config.app_version.clone(),
|
||||||
)?;
|
)?;
|
||||||
let harness = control.cursor_harness().clone();
|
let harness = control.cursor_harness().clone();
|
||||||
let mut router = api::router(registry.clone(), clients)?;
|
let mut router = api::router(registry.clone(), clients)?;
|
||||||
|
|||||||
@@ -16,7 +16,8 @@ use super::ControlService;
|
|||||||
// 此广告拉取不涉及用户隐私,用户id随机产生
|
// 此广告拉取不涉及用户隐私,用户id随机产生
|
||||||
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
|
// 开源项目广告为作者唯一收入来源,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 DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
|
||||||
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
||||||
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ pub struct ControlService {
|
|||||||
plugin_runtime: PluginRuntime,
|
plugin_runtime: PluginRuntime,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
clients: crate::network::NetworkClients,
|
clients: crate::network::NetworkClients,
|
||||||
|
app_version: String,
|
||||||
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -153,6 +154,7 @@ impl ControlService {
|
|||||||
plugin_runtime: PluginRuntime,
|
plugin_runtime: PluginRuntime,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
clients: crate::network::NetworkClients,
|
clients: crate::network::NetworkClients,
|
||||||
|
app_version: String,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
cursor_harness: CursorHarness::new(store.clone())?,
|
cursor_harness: CursorHarness::new(store.clone())?,
|
||||||
@@ -161,6 +163,7 @@ impl ControlService {
|
|||||||
plugin_runtime,
|
plugin_runtime,
|
||||||
plugins,
|
plugins,
|
||||||
clients,
|
clients,
|
||||||
|
app_version,
|
||||||
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -265,7 +268,7 @@ impl ControlService {
|
|||||||
.get(ADS_ENDPOINT)
|
.get(ADS_ENDPOINT)
|
||||||
.header(DEVICE_ID_HEADER, installation_id)
|
.header(DEVICE_ID_HEADER, installation_id)
|
||||||
.header(OS_HEADER, std::env::consts::OS)
|
.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)
|
.header(LANGUAGE_HEADER, language)
|
||||||
.timeout(std::time::Duration::from_secs(60));
|
.timeout(std::time::Duration::from_secs(60));
|
||||||
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
|
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
|
||||||
@@ -299,7 +302,7 @@ impl ControlService {
|
|||||||
.post(endpoint)
|
.post(endpoint)
|
||||||
.header(DEVICE_ID_HEADER, installation_id)
|
.header(DEVICE_ID_HEADER, installation_id)
|
||||||
.header(OS_HEADER, std::env::consts::OS)
|
.header(OS_HEADER, std::env::consts::OS)
|
||||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
.header(APP_VERSION_HEADER, &self.app_version)
|
||||||
.json(input)
|
.json(input)
|
||||||
.timeout(std::time::Duration::from_secs(5))
|
.timeout(std::time::Duration::from_secs(5))
|
||||||
.send()
|
.send()
|
||||||
|
|||||||
@@ -2,6 +2,12 @@
|
|||||||
|
|
||||||
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
|
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
|
||||||
|
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub enum RunFinish {
|
||||||
|
TurnCompleted,
|
||||||
|
Transport(TransportFinish),
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub enum TransportFinish {
|
pub enum TransportFinish {
|
||||||
Success,
|
Success,
|
||||||
@@ -17,7 +23,7 @@ pub enum TransportCommand {
|
|||||||
},
|
},
|
||||||
RunFinished {
|
RunFinished {
|
||||||
generation: u64,
|
generation: u64,
|
||||||
finish: TransportFinish,
|
finish: RunFinish,
|
||||||
},
|
},
|
||||||
Disconnect,
|
Disconnect,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ use crate::{
|
|||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, TransportFinish};
|
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
|
||||||
use crate::cursor::transport::TransportHandle;
|
use crate::cursor::transport::TransportHandle;
|
||||||
|
|
||||||
pub struct ConversationOutput {
|
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;
|
let result = self.run_inner().await;
|
||||||
if let Err(error) = &result {
|
if let Err(error) = &result {
|
||||||
if !self.superseded.is_cancelled() {
|
if !self.superseded.is_cancelled() {
|
||||||
@@ -143,7 +143,7 @@ impl ConversationOutput {
|
|||||||
result
|
result
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn run_inner(&mut self) -> Result<TransportFinish> {
|
async fn run_inner(&mut self) -> Result<RunFinish> {
|
||||||
if self.context.compacting {
|
if self.context.compacting {
|
||||||
self.handle.emit(&events::summary_started())?;
|
self.handle.emit(&events::summary_started())?;
|
||||||
}
|
}
|
||||||
@@ -175,7 +175,7 @@ impl ConversationOutput {
|
|||||||
if self.superseded.is_cancelled() {
|
if self.superseded.is_cancelled() {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
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() {
|
let input = if let Ok(action) = self.runtime_actions.try_recv() {
|
||||||
Input::RuntimeAction(Some(Box::new(action)))
|
Input::RuntimeAction(Some(Box::new(action)))
|
||||||
@@ -187,7 +187,7 @@ impl ConversationOutput {
|
|||||||
_ = self.superseded.cancelled() => {
|
_ = self.superseded.cancelled() => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
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)),
|
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
||||||
event = self.core.events.recv() => Input::Event(event),
|
event = self.core.events.recv() => Input::Event(event),
|
||||||
@@ -719,7 +719,7 @@ impl ConversationOutput {
|
|||||||
if self.superseded.is_cancelled() {
|
if self.superseded.is_cancelled() {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
return Ok(TransportFinish::Cancelled);
|
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||||
}
|
}
|
||||||
return match outcome {
|
return match outcome {
|
||||||
RunOutcome::Completed => {
|
RunOutcome::Completed => {
|
||||||
@@ -735,7 +735,7 @@ impl ConversationOutput {
|
|||||||
for _ in 0..3 {
|
for _ in 0..3 {
|
||||||
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
||||||
}
|
}
|
||||||
return Ok(TransportFinish::Success);
|
return Ok(RunFinish::TurnCompleted);
|
||||||
}
|
}
|
||||||
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
||||||
Error::Protocol("Completed without final state".into())
|
Error::Protocol("Completed without final state".into())
|
||||||
@@ -751,17 +751,19 @@ impl ConversationOutput {
|
|||||||
ttft_breakdown: None,
|
ttft_breakdown: None,
|
||||||
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
||||||
})?;
|
})?;
|
||||||
Ok(TransportFinish::Success)
|
Ok(RunFinish::TurnCompleted)
|
||||||
}
|
}
|
||||||
RunOutcome::Cancelled => {
|
RunOutcome::Cancelled => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
Ok(TransportFinish::Cancelled)
|
Ok(RunFinish::Transport(TransportFinish::Cancelled))
|
||||||
}
|
}
|
||||||
RunOutcome::Failed(failure) => {
|
RunOutcome::Failed(failure) => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
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},
|
transport::{OrderedInbox, TransportHandle},
|
||||||
},
|
},
|
||||||
run::{CommandResult, RunEngine, RunHandle},
|
run::{CommandResult, RunEngine, RunHandle, RunPhase},
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
||||||
ConversationRegistry, MessageDelivery, TransportCommand, TransportFinish,
|
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct ConversationRuntime;
|
pub struct ConversationRuntime;
|
||||||
|
|
||||||
|
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct RunGeneration {
|
struct RunGeneration {
|
||||||
id: u64,
|
id: u64,
|
||||||
|
request: pb::AgentRunRequest,
|
||||||
superseded: CancellationToken,
|
superseded: CancellationToken,
|
||||||
finished: CancellationToken,
|
finished: CancellationToken,
|
||||||
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
||||||
@@ -78,6 +81,7 @@ impl ConversationRuntime {
|
|||||||
let mut next_generation = 1_u64;
|
let mut next_generation = 1_u64;
|
||||||
let mut pending_finish = None::<(u64, TransportFinish)>;
|
let mut pending_finish = None::<(u64, TransportFinish)>;
|
||||||
let mut draining = false;
|
let mut draining = false;
|
||||||
|
let mut waiting_for_action = false;
|
||||||
loop {
|
loop {
|
||||||
let command = if draining {
|
let command = if draining {
|
||||||
if !handle.admissions_drained() {
|
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 {
|
} else {
|
||||||
match receiver.recv().await {
|
match receiver.recv().await {
|
||||||
Some(command) => command,
|
Some(command) => command,
|
||||||
@@ -130,7 +156,19 @@ impl ConversationRuntime {
|
|||||||
let _ = handle.emit(&codec::abort(id));
|
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();
|
super::finish_cancelled(&handle).ok();
|
||||||
|
}
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
TransportCommand::RunFinished { generation, finish } => {
|
TransportCommand::RunFinished { generation, finish } => {
|
||||||
@@ -140,10 +178,20 @@ impl ConversationRuntime {
|
|||||||
{
|
{
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
match finish {
|
||||||
|
RunFinish::TurnCompleted => {
|
||||||
|
pending_finish = None;
|
||||||
|
draining = false;
|
||||||
|
waiting_for_action = true;
|
||||||
|
}
|
||||||
|
RunFinish::Transport(finish) => {
|
||||||
|
waiting_for_action = false;
|
||||||
handle.begin_close();
|
handle.begin_close();
|
||||||
pending_finish = Some((generation, finish));
|
pending_finish = Some((generation, finish));
|
||||||
draining = true;
|
draining = true;
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
TransportCommand::Append { seqno, message } => {
|
TransportCommand::Append { seqno, message } => {
|
||||||
for (_seqno, message) in inbox.push(seqno, *message) {
|
for (_seqno, message) in inbox.push(seqno, *message) {
|
||||||
{
|
{
|
||||||
@@ -151,6 +199,7 @@ impl ConversationRuntime {
|
|||||||
Some(pb::agent_client_message::Message::RunRequest(
|
Some(pb::agent_client_message::Message::RunRequest(
|
||||||
request,
|
request,
|
||||||
)) => {
|
)) => {
|
||||||
|
waiting_for_action = false;
|
||||||
if draining {
|
if draining {
|
||||||
handle.reopen();
|
handle.reopen();
|
||||||
draining = false;
|
draining = false;
|
||||||
@@ -171,57 +220,18 @@ impl ConversationRuntime {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
let previous_finished =
|
start_generation(
|
||||||
if let Some(previous) = current.take() {
|
®istry,
|
||||||
previous.superseded.cancel();
|
&handle,
|
||||||
if let Some(run) = previous.run.lock().clone() {
|
&dependencies,
|
||||||
run.cancel();
|
&blob_sync,
|
||||||
}
|
&context_sync,
|
||||||
for id in previous
|
&tool_runtime_factory,
|
||||||
.tool_runtime
|
&mut current,
|
||||||
.interrupt_for_run_replacement()
|
&mut next_generation,
|
||||||
.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(),
|
|
||||||
request,
|
request,
|
||||||
dependencies.clone(),
|
)
|
||||||
blob_sync.clone(),
|
.await;
|
||||||
context_sync.clone(),
|
|
||||||
generation,
|
|
||||||
previous_finished,
|
|
||||||
result_receiver,
|
|
||||||
runtime_action_receiver,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||||
message,
|
message,
|
||||||
@@ -382,27 +392,48 @@ impl ConversationRuntime {
|
|||||||
// return an explicit Protocol Error rather than falling through silently.
|
// return an explicit Protocol Error rather than falling through silently.
|
||||||
Some(
|
Some(
|
||||||
pb::agent_client_message::Message::ConversationAction(
|
pb::agent_client_message::Message::ConversationAction(
|
||||||
action,
|
conversation_action,
|
||||||
),
|
),
|
||||||
) => match action.action {
|
) => match conversation_action.action.clone() {
|
||||||
Some(
|
Some(
|
||||||
pb::conversation_action::Action::UserMessageAction(
|
pb::conversation_action::Action::UserMessageAction(
|
||||||
action,
|
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;
|
continue;
|
||||||
};
|
};
|
||||||
if generation
|
let mut request = previous.request.clone();
|
||||||
.runtime_actions
|
request.action = Some(conversation_action);
|
||||||
.send(compile::RuntimeAction::UserMessage(action))
|
request.conversation_state = None;
|
||||||
.is_err()
|
request.pre_fetched_blobs.clear();
|
||||||
{
|
waiting_for_action = false;
|
||||||
generation.results.send_error(crate::Error::Protocol(
|
start_generation(
|
||||||
"UserMessageAction arrived without an active Run"
|
®istry,
|
||||||
.into(),
|
&handle,
|
||||||
));
|
&dependencies,
|
||||||
}
|
&blob_sync,
|
||||||
|
&context_sync,
|
||||||
|
&tool_runtime_factory,
|
||||||
|
&mut current,
|
||||||
|
&mut next_generation,
|
||||||
|
request,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||||
if let Some(generation) = current.as_ref() {
|
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(
|
fn finish_pending(
|
||||||
handle: &TransportHandle,
|
handle: &TransportHandle,
|
||||||
current: &Option<RunGeneration>,
|
current: &Option<RunGeneration>,
|
||||||
@@ -569,7 +661,7 @@ fn spawn_run_request(
|
|||||||
let _ = handle
|
let _ = handle
|
||||||
.command(TransportCommand::RunFinished {
|
.command(TransportCommand::RunFinished {
|
||||||
generation: generation.id,
|
generation: generation.id,
|
||||||
finish: TransportFinish::Failed(error),
|
finish: RunFinish::Transport(TransportFinish::Failed(error)),
|
||||||
})
|
})
|
||||||
.await;
|
.await;
|
||||||
return;
|
return;
|
||||||
@@ -606,7 +698,7 @@ fn spawn_run_request(
|
|||||||
let _ = handle
|
let _ = handle
|
||||||
.command(TransportCommand::RunFinished {
|
.command(TransportCommand::RunFinished {
|
||||||
generation: generation.id,
|
generation: generation.id,
|
||||||
finish: TransportFinish::Success,
|
finish: RunFinish::Transport(TransportFinish::Success),
|
||||||
})
|
})
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
@@ -635,7 +727,7 @@ fn spawn_run_request(
|
|||||||
let _ = handle
|
let _ = handle
|
||||||
.command(TransportCommand::RunFinished {
|
.command(TransportCommand::RunFinished {
|
||||||
generation: generation.id,
|
generation: generation.id,
|
||||||
finish: TransportFinish::Success,
|
finish: RunFinish::Transport(TransportFinish::Success),
|
||||||
})
|
})
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
@@ -700,14 +792,14 @@ fn spawn_run_request(
|
|||||||
Ok(finish) => finish,
|
Ok(finish) => finish,
|
||||||
Err(error) => {
|
Err(error) => {
|
||||||
if generation.superseded.is_cancelled() {
|
if generation.superseded.is_cancelled() {
|
||||||
TransportFinish::Cancelled
|
RunFinish::Transport(TransportFinish::Cancelled)
|
||||||
} else {
|
} else {
|
||||||
tracing::error!(
|
tracing::error!(
|
||||||
request_id = handle.request_id(),
|
request_id = handle.request_id(),
|
||||||
%error,
|
%error,
|
||||||
"Cursor session failed"
|
"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);
|
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]
|
#[tokio::test]
|
||||||
async fn runtime_user_message_action_interrupts_and_continues_with_new_message() {
|
async fn runtime_user_message_action_interrupts_and_continues_with_new_message() {
|
||||||
let (_directory, store) = fixtures::temp_store().await;
|
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(
|
async fn acknowledge_kv(
|
||||||
handle: &cursor_server::cursor::TransportHandle,
|
handle: &cursor_server::cursor::TransportHandle,
|
||||||
append_seqno: &mut i64,
|
append_seqno: &mut i64,
|
||||||
|
|||||||
Reference in New Issue
Block a user