diff --git a/server/src/app.rs b/server/src/app.rs index 62cb015..6f97c86 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -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)?; diff --git a/server/src/control/ads.rs b/server/src/control/ads.rs index 7749f65..bd5ffe9 100644 --- a/server/src/control/ads.rs +++ b/server/src/control/ads.rs @@ -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"; diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 28e9e09..7a9d827 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -42,6 +42,7 @@ pub struct ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, model_tests: Arc>>, } @@ -153,6 +154,7 @@ impl ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, ) -> Result { 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() diff --git a/server/src/cursor/conversation/command.rs b/server/src/cursor/conversation/command.rs index fa97fe4..4ba5f2b 100644 --- a/server/src/cursor/conversation/command.rs +++ b/server/src/cursor/conversation/command.rs @@ -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, } diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index 86765e8..349fccc 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -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 { + pub async fn run(mut self) -> Result { 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 { + async fn run_inner(&mut self) -> Result { 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, + )))) } }; } diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index 4f6d3f6..3a1edec 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -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>>, @@ -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)); } } - super::finish_cancelled(&handle).ok(); + 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,9 +178,19 @@ impl ConversationRuntime { { continue; } - handle.begin_close(); - pending_finish = Some((generation, finish)); - draining = true; + 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::(); - 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, + 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::(); + 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, @@ -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)) } } }; diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 1cfd882..9c50c8f 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -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, + 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, + 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,