diff --git a/apps/desktop/src/api.ts b/apps/desktop/src/api.ts index 4a04b26..13e1f1c 100644 --- a/apps/desktop/src/api.ts +++ b/apps/desktop/src/api.ts @@ -86,7 +86,7 @@ export interface LegacyModelImportResult { export interface ModelConnectivityResult { duration_ms: number; - first_text_ms: number | null; + first_valid_response_ms: number | null; output_tokens: number; tokens_per_second: number; tokens_estimated: boolean; @@ -197,6 +197,7 @@ export interface LlmCall { created_at_ms: number; ttfb_ms: number | null; ttft_ms: number | null; + ttfr_ms: number | null; duration_ms: number | null; input_tokens: number | null; output_tokens: number | null; diff --git a/apps/desktop/src/components/CallDetails.tsx b/apps/desktop/src/components/CallDetails.tsx index 452306d..e365ae2 100644 --- a/apps/desktop/src/components/CallDetails.tsx +++ b/apps/desktop/src/components/CallDetails.tsx @@ -32,6 +32,7 @@ export function CallDetails({ detail }: { detail: CallDetail }) { ["Created At", `${call.created_at_ms} · ${new Date(call.created_at_ms).toLocaleString()}`], [t("耗时"), timing(call.duration_ms)], ["TTFB", timing(call.ttfb_ms)], + ["TTFR", timing(call.ttfr_ms)], ["TTFT", timing(call.ttft_ms)], ["Input Token", show(call.input_tokens)], ["Output Token", show(call.output_tokens)], diff --git a/apps/desktop/src/components/CallTable.tsx b/apps/desktop/src/components/CallTable.tsx index ac7fb74..ee304e3 100644 --- a/apps/desktop/src/components/CallTable.tsx +++ b/apps/desktop/src/components/CallTable.tsx @@ -87,6 +87,11 @@ export function CallTable({ calls, onDetails }: { calls: LlmCall[]; onDetails: ( header: "TTFB", render: (call) => milliseconds(call.ttfb_ms), }, + { + key: "ttfr", + header: "TTFR", + render: (call) => milliseconds(call.ttfr_ms), + }, { key: "ttft", header: "TTFT", diff --git a/apps/desktop/src/components/charts/LatencyChart.tsx b/apps/desktop/src/components/charts/LatencyChart.tsx index 502b0fc..27886ac 100644 --- a/apps/desktop/src/components/charts/LatencyChart.tsx +++ b/apps/desktop/src/components/charts/LatencyChart.tsx @@ -6,13 +6,14 @@ export function LatencyChart({ calls }: { calls: LlmCall[] }) { const points = calls.slice(0, 20).reverse(); const option: EChartsCoreOption = { animationDuration: 280, - color: ["#d79a62", "#ad84cf"], + color: ["#d79a62", "#ad84cf", "#72a8d8"], grid: { top: 22, right: 18, bottom: 30, left: 48 }, legend: { top: 0, right: 8, textStyle: { color: "#999" } }, tooltip: { trigger: "axis" }, xAxis: { type: "category", data: points.map((call) => new Date(call.created_at_ms).toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" })), axisLabel: { color: "#888" }, axisLine: { lineStyle: { color: "#5555" } } }, yAxis: { type: "value", axisLabel: { color: "#888", formatter: "{value} ms" }, splitLine: { lineStyle: { color: "#8882" } } }, series: [ + { name: "TTFR", type: "line", smooth: true, symbol: "none", data: points.map((call) => call.ttfr_ms ?? 0) }, { name: "TTFT", type: "line", smooth: true, symbol: "none", data: points.map((call) => call.ttft_ms ?? 0) }, { name: t("总耗时"), type: "line", smooth: true, symbol: "none", data: points.map((call) => call.duration_ms ?? 0) }, ], diff --git a/apps/desktop/src/components/cursor/CursorModelTestResult.tsx b/apps/desktop/src/components/cursor/CursorModelTestResult.tsx index e6f065b..fac62b5 100644 --- a/apps/desktop/src/components/cursor/CursorModelTestResult.tsx +++ b/apps/desktop/src/components/cursor/CursorModelTestResult.tsx @@ -21,7 +21,7 @@ export function CursorModelTestResult({ state, testing = false }: { state?: Curs const detail = success ? t("速度 {speed} tokens/s · 首字 {firstText} ms · 总耗时 {duration} ms · 输出 {tokens} tokens{estimated} · 返回:{output}", { speed: formatSpeed(state.result.tokens_per_second), - firstText: state.result.first_text_ms ?? "--", + firstText: state.result.first_valid_response_ms ?? "--", duration: state.result.duration_ms, tokens: state.result.output_tokens, estimated: state.result.tokens_estimated ? t("(估算)") : "", diff --git a/apps/desktop/src/demo/api.ts b/apps/desktop/src/demo/api.ts index 38bf1f3..4500617 100644 --- a/apps/desktop/src/demo/api.ts +++ b/apps/desktop/src/demo/api.ts @@ -55,6 +55,7 @@ const calls: LlmCall[] = Array.from({ length: 24 }, (_, index) => { finish_reason: failed ? null : "stop", created_at_ms: FIXED_NOW - index * 3 * 60_000, ttfb_ms: 210 + index * 13, + ttfr_ms: 290 + index * 15, ttft_ms: 370 + index * 17, duration_ms: failed ? 812 : 1_420 + index * 71, input_tokens: 4_800 + index * 337, @@ -116,7 +117,7 @@ export function installDemoApi() { } if (path === "/models/import-v0049") return json({ imported: 0, skipped: 0, total: 0 }); if (/^\/models\/[^/]+\/test\/[^/]+$/.test(path) && method === "POST") { - return json({ duration_ms: 1_284, first_text_ms: 418, output_tokens: 42, tokens_per_second: 38.6, tokens_estimated: false, output: "Mock connectivity test passed." }); + return json({ duration_ms: 1_284, first_valid_response_ms: 418, output_tokens: 42, tokens_per_second: 38.6, tokens_estimated: false, output: "Mock connectivity test passed." }); } if (/^\/models\/[^/]+\/test\/[^/]+$/.test(path) || /^\/models\/[^/]+$/.test(path)) { return method === "DELETE" ? empty() : json(models[0]); diff --git a/server/migrations/0005_add_first_valid_response_timing.sql b/server/migrations/0005_add_first_valid_response_timing.sql new file mode 100644 index 0000000..981defc --- /dev/null +++ b/server/migrations/0005_add_first_valid_response_timing.sql @@ -0,0 +1,2 @@ +ALTER TABLE llm_calls ADD COLUMN first_valid_response_at_ms INTEGER; +ALTER TABLE llm_calls ADD COLUMN ttfr_ms INTEGER; diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 3e9a7b0..1f6ff10 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -24,7 +24,7 @@ use crate::{ ModelRequest, ModelSpec, ModelType, Overview, ProjectedContent, ProjectedMessage, PromptSpec, ProviderType, Role, }, - provider::{ModelEvent, Provider}, + provider::{is_valid_response_event, ModelEvent, Provider}, store::{ DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store, TabSettings, @@ -95,7 +95,7 @@ fn empty_json_object_ref() -> &'static serde_json::Value { #[derive(Clone, Debug, Serialize)] pub struct ModelConnectivityResult { pub duration_ms: u64, - pub first_text_ms: Option, + pub first_valid_response_ms: Option, pub output_tokens: u64, pub tokens_per_second: f64, pub tokens_estimated: bool, @@ -322,7 +322,7 @@ impl ControlService { }, }; let started = Instant::now(); - let mut first_text_at = None; + let mut first_valid_response_at = None; let mut output_tokens = None; let mut output = String::new(); let stream = self.provider.stream(invocation, cancellation.clone()); @@ -330,11 +330,12 @@ impl ControlService { futures_util::pin_mut!(stream); let mut finished = false; while let Some(event) = stream.next().await { - match event? { + let event = event?; + if first_valid_response_at.is_none() && is_valid_response_event(&event) { + first_valid_response_at = Some(Instant::now()); + } + match event { ModelEvent::TextDelta(delta) => { - if first_text_at.is_none() && !delta.trim().is_empty() { - first_text_at = Some(Instant::now()); - } output.push_str(&delta); } ModelEvent::Usage(usage) => { @@ -380,16 +381,16 @@ impl ControlService { } let elapsed = started.elapsed(); let output = output.trim().to_string(); - if first_text_at.is_none() { + if first_valid_response_at.is_none() { return Err(Error::Provider( - "model connectivity test received no text output".into(), + "model connectivity test received no valid response".into(), )); } let tokens_estimated = output_tokens.is_none(); let output_tokens = output_tokens.unwrap_or_else(|| estimate_output_tokens(&output)); Ok(ModelConnectivityResult { duration_ms: elapsed.as_millis().min(u128::from(u64::MAX)) as u64, - first_text_ms: first_text_at.map(|first| { + first_valid_response_ms: first_valid_response_at.map(|first| { first .duration_since(started) .as_millis() @@ -628,10 +629,12 @@ fn official_call(trace: CursorRunTraceSummary) -> CallSummary { response_headers_at_ms: trace.first_response_at_ms, first_event_at_ms: trace.first_response_at_ms, first_text_at_ms: None, + first_valid_response_at_ms: None, finished_at_ms: trace.finished_at_ms, queue_ms: None, ttfb_ms: ttfb, ttft_ms: None, + ttfr_ms: None, duration_ms: duration, input_tokens: None, output_tokens: None, @@ -1132,6 +1135,7 @@ mod tests { .unwrap(); assert_eq!(result.output, "OK"); + assert!(result.first_valid_response_ms.is_some()); assert_eq!(result.output_tokens, 2); assert!(!result.tokens_estimated); assert!(result.tokens_per_second > 0.0); diff --git a/server/src/cursor/actor.rs b/server/src/cursor/actor.rs index b8e31bf..9cfe5a0 100644 --- a/server/src/cursor/actor.rs +++ b/server/src/cursor/actor.rs @@ -44,7 +44,8 @@ impl CursorActor { tokio::spawn(async move { let mut inbox = OrderedInbox::starting_at(next_append_seqno); let (results_tx, results_rx) = tool_result_channel(); - let (runtime_actions_tx, runtime_actions_rx) = mpsc::unbounded_channel(); + let (runtime_actions_tx, runtime_actions_rx) = + mpsc::unbounded_channel::(); let tool_runtime = CursorToolRuntime::default(); let context_sync = RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone()); @@ -375,10 +376,21 @@ impl CursorActor { ), ) => match action.action { Some( - pb::conversation_action::Action::UserMessageAction(_), + pb::conversation_action::Action::UserMessageAction( + action, + ), ) => { - handle.mark_conversation_cancelled(); - handle.cancel(); + if runtime_actions_tx + .send(super::request::RuntimeAction::UserMessage( + action, + )) + .is_err() + { + results_tx.send_error(crate::Error::Protocol( + "UserMessageAction arrived without an active Run" + .into(), + )); + } } Some(pb::conversation_action::Action::CancelAction(_)) => { handle.mark_conversation_cancelled(); @@ -389,7 +401,10 @@ impl CursorActor { action, ), ) => { - if runtime_actions_tx.send(action).is_err() { + if runtime_actions_tx + .send(super::request::RuntimeAction::Inject(action)) + .is_err() + { results_tx.send_error(crate::Error::Protocol( "InjectContextAction arrived without an active Run" .into(), diff --git a/server/src/cursor/bidi_append.rs b/server/src/cursor/bidi_append.rs index 73c55d9..bcfc54b 100644 --- a/server/src/cursor/bidi_append.rs +++ b/server/src/cursor/bidi_append.rs @@ -65,7 +65,6 @@ impl DecodedAppend { if matches!( action.action.as_ref(), Some(agent::conversation_action::Action::CancelAction(_)) - | Some(agent::conversation_action::Action::UserMessageAction(_)) ) ) } diff --git a/server/src/cursor/request/mod.rs b/server/src/cursor/request/mod.rs index a5bed26..ec9463b 100644 --- a/server/src/cursor/request/mod.rs +++ b/server/src/cursor/request/mod.rs @@ -6,4 +6,4 @@ mod prepare; mod runtime; pub use prepare::*; -pub(crate) use runtime::compile_injection; +pub(crate) use runtime::{compile_injection, compile_user_message_action, RuntimeAction}; diff --git a/server/src/cursor/request/runtime.rs b/server/src/cursor/request/runtime.rs index 0d3583c..929fc23 100644 --- a/server/src/cursor/request/runtime.rs +++ b/server/src/cursor/request/runtime.rs @@ -16,6 +16,57 @@ use crate::{ use super::{context, images}; +pub(crate) enum RuntimeAction { + Inject(pb::InjectContextAction), + UserMessage(pb::UserMessageAction), +} + +pub(crate) async fn compile_user_message_action( + action: &pb::UserMessageAction, + current_mode: i32, + compiler: &PromptCompiler, + blobs: &BlobSynchronizer, +) -> Result { + let user = action + .user_message + .as_ref() + .ok_or_else(|| Error::Protocol("Cursor user message action has no UserMessage".into()))?; + if user.message_id.is_empty() { + return Err(Error::Protocol( + "Cursor user message action has no message_id".into(), + )); + } + let mode = if user.mode == pb::AgentMode::Unspecified as i32 { + current_mode + } else { + user.mode + }; + let mut action_context = action + .prepend_user_messages + .iter() + .map(|message| message.text.trim()) + .filter(|text| !text.is_empty()) + .map(str::to_string) + .collect::>(); + action_context.extend( + user.subagent_system_reminder + .iter() + .filter(|text| !text.is_empty()) + .cloned(), + ); + let empty_context = pb::RequestContext::default(); + compile( + format!("user-message:{}", user.message_id), + super::prepare::mode_from_proto(mode)?, + user, + action.request_context.as_ref().unwrap_or(&empty_context), + &action_context.join("\\n\\n"), + compiler, + blobs, + ) + .await +} + pub(crate) async fn compile_injection( injection: &pb::InjectContextAction, mode: i32, diff --git a/server/src/cursor/session.rs b/server/src/cursor/session.rs index 2999f75..778c9a3 100644 --- a/server/src/cursor/session.rs +++ b/server/src/cursor/session.rs @@ -14,7 +14,9 @@ use crate::{ presentation::Presentation, prompting::PromptCompiler, proto::agent::v1 as pb, - request::CursorRunContext, + request::{ + compile_injection, compile_user_message_action, CursorRunContext, RuntimeAction, + }, tools::{ codec, result::{ToolCompletion, ToolResultReceiver}, @@ -40,7 +42,7 @@ pub struct CursorSession { results: ToolResultReceiver, checkpoint: CheckpointBuilder, tool_runtime: CursorToolRuntime, - runtime_actions: mpsc::UnboundedReceiver, + runtime_actions: mpsc::UnboundedReceiver, compiler: PromptCompiler, blob_sync: BlobSynchronizer, injection_ids: HashSet, @@ -52,12 +54,20 @@ struct PendingInjection { delivery_batch_id: String, } +struct InjectionState<'a> { + active_round: Option<&'a ToolRoundId>, + active_tool_calls: &'a HashSet, + completions: &'a HashMap, + interrupted_rounds: &'a mut HashSet, + interrupted_tool_calls: &'a mut HashSet, +} + pub(crate) struct CursorSessionRuntime { pub tools: ToolDispatcher, pub results: ToolResultReceiver, pub checkpoint: CheckpointBuilder, pub tool_runtime: CursorToolRuntime, - pub runtime_actions: mpsc::UnboundedReceiver, + pub runtime_actions: mpsc::UnboundedReceiver, pub compiler: PromptCompiler, pub blob_sync: BlobSynchronizer, } @@ -170,17 +180,30 @@ impl CursorSession { Input::CompletionResult(None) => { return Err(Error::Protocol("tool result channel closed".into())); } - Input::RuntimeAction(Some(action)) => { - self.forward_injection( - *action, - active_round.as_ref(), - &active_tool_calls, - &completions, - &mut interrupted_rounds, - &mut interrupted_tool_calls, - ) - .await?; - } + Input::RuntimeAction(Some(action)) => match *action { + RuntimeAction::Inject(action) => { + self.forward_injection( + action, + active_round.as_ref(), + &active_tool_calls, + &completions, + &mut interrupted_rounds, + &mut interrupted_tool_calls, + ) + .await?; + } + RuntimeAction::UserMessage(action) => { + self.forward_user_message( + action, + active_round.as_ref(), + &active_tool_calls, + &completions, + &mut interrupted_rounds, + &mut interrupted_tool_calls, + ) + .await?; + } + }, Input::RuntimeAction(None) => { return Err(Error::Protocol("runtime action channel closed".into())); } @@ -679,6 +702,41 @@ impl CursorSession { Ok(dispatched.completion) } + async fn forward_user_message( + &mut self, + action: pb::UserMessageAction, + active_round: Option<&ToolRoundId>, + active_tool_calls: &HashSet, + completions: &HashMap, + interrupted_rounds: &mut HashSet, + interrupted_tool_calls: &mut HashSet, + ) -> Result<()> { + let user_message = action.user_message.clone().ok_or_else(|| { + Error::Protocol("Cursor user message action has no UserMessage".into()) + })?; + let injection_id = format!("user-message:{}", user_message.message_id); + let message = compile_user_message_action( + &action, + self.context.mode, + &self.compiler, + &self.blob_sync, + ) + .await?; + self.queue_injection( + injection_id, + Some(user_message), + message, + InjectionState { + active_round, + active_tool_calls, + completions, + interrupted_rounds, + interrupted_tool_calls, + }, + ) + .await + } + async fn forward_injection( &mut self, action: pb::InjectContextAction, @@ -714,14 +772,30 @@ impl CursorSession { } _ => None, }; - let message = crate::cursor::request::compile_injection( - &action, - self.context.mode, - &self.compiler, - &self.blob_sync, + let message = + compile_injection(&action, self.context.mode, &self.compiler, &self.blob_sync).await?; + self.queue_injection( + action.injection_id, + user_message, + message, + InjectionState { + active_round, + active_tool_calls, + completions, + interrupted_rounds, + interrupted_tool_calls, + }, ) - .await?; - let injection_id = action.injection_id; + .await + } + + async fn queue_injection( + &mut self, + injection_id: String, + user_message: Option, + message: crate::model::CanonicalMessage, + state: InjectionState<'_>, + ) -> Result<()> { let delivery_batch_id = injection_id.clone(); self.injection_ids.insert(injection_id.clone()); self.pending_injections.insert( @@ -733,14 +807,15 @@ impl CursorSession { ); self.handle .emit(&interaction::context_injection_queued(injection_id.clone()))?; - interrupted_tool_calls.extend( - active_tool_calls + state.interrupted_tool_calls.extend( + state + .active_tool_calls .iter() - .filter(|call_id| !completions.contains_key(*call_id)) + .filter(|call_id| !state.completions.contains_key(*call_id)) .cloned(), ); - if let Some(round_id) = active_round { - interrupted_rounds.insert(round_id.clone()); + if let Some(round_id) = state.active_round { + state.interrupted_rounds.insert(round_id.clone()); } self.interrupt_execs().await; if self @@ -780,7 +855,7 @@ enum Input { Event(Option), Completion(ToolCompletion), CompletionResult(Option>), - RuntimeAction(Option>), + RuntimeAction(Option>), CheckpointFailure(Option), } diff --git a/server/src/model/llm_call.rs b/server/src/model/llm_call.rs index 123a80b..8cb636f 100644 --- a/server/src/model/llm_call.rs +++ b/server/src/model/llm_call.rs @@ -52,10 +52,12 @@ pub struct LlmCallSummary { pub response_headers_at_ms: Option, pub first_event_at_ms: Option, pub first_text_at_ms: Option, + pub first_valid_response_at_ms: Option, pub finished_at_ms: Option, pub queue_ms: Option, pub ttfb_ms: Option, pub ttft_ms: Option, + pub ttfr_ms: Option, pub duration_ms: Option, pub input_tokens: Option, pub output_tokens: Option, diff --git a/server/src/provider/event.rs b/server/src/provider/event.rs index 861cafd..ce512e8 100644 --- a/server/src/provider/event.rs +++ b/server/src/provider/event.rs @@ -34,3 +34,56 @@ pub enum ModelEvent { Usage(Usage), Done(FinishReason), } + +/// Returns whether an event represents the first valid upstream response. +/// Transport markers, replay metadata, usage, completion, and provider heartbeats +/// are intentionally excluded; empty text/reasoning/tool deltas are valid events. +pub fn is_valid_response_event(event: &ModelEvent) -> bool { + matches!( + event, + ModelEvent::TextDelta(_) + | ModelEvent::ThinkingStart + | ModelEvent::ThinkingDelta(_) + | ModelEvent::ThinkingEnd + | ModelEvent::ToolCallStart { .. } + | ModelEvent::ToolCallArgumentsDelta { .. } + | ModelEvent::ToolCallEnd { .. } + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn response_markers_include_empty_content_but_exclude_transport_events() { + assert!(is_valid_response_event(&ModelEvent::TextDelta( + String::new() + ))); + assert!(is_valid_response_event(&ModelEvent::ThinkingDelta( + String::new() + ))); + assert!(is_valid_response_event( + &ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: String::new(), + } + )); + assert!(is_valid_response_event(&ModelEvent::ThinkingStart)); + assert!(is_valid_response_event(&ModelEvent::ToolCallStart { + index: 0, + call_id: "call".into(), + name: "tool".into(), + })); + assert!(!is_valid_response_event(&ModelEvent::Start { + model_call_id: "call".into(), + })); + assert!(!is_valid_response_event(&ModelEvent::TextStart)); + assert!(!is_valid_response_event(&ModelEvent::Usage( + Usage::default() + ))); + assert!(!is_valid_response_event(&ModelEvent::Done( + FinishReason::Stop + ))); + } +} diff --git a/server/src/provider/recorder.rs b/server/src/provider/recorder.rs index b65169f..29aa0cc 100644 --- a/server/src/provider/recorder.rs +++ b/server/src/provider/recorder.rs @@ -14,7 +14,7 @@ use crate::{ Result, }; -use super::{FinishReason, ModelEvent}; +use super::{is_valid_response_event, FinishReason, ModelEvent}; pub(crate) fn recorded_headers( config: &crate::config::ProviderConfig, @@ -56,6 +56,7 @@ struct AttemptState { next_chunk: AtomicI64, chunks: ChunkBuffer, first_text_recorded: AtomicBool, + first_valid_response_recorded: AtomicBool, } impl AttemptState { @@ -66,6 +67,7 @@ impl AttemptState { next_chunk: AtomicI64::new(0), chunks: ChunkBuffer::default(), first_text_recorded: AtomicBool::new(false), + first_valid_response_recorded: AtomicBool::new(false), } } } @@ -176,8 +178,29 @@ impl CallRecorder { } pub async fn event(&self, event: &ModelEvent) -> Result<()> { + let attempt = self.inner.attempt.lock().await; + if is_valid_response_event(event) + && attempt + .first_valid_response_recorded + .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + if let Err(error) = self + .inner + .store + .record_llm_first_valid_response(&attempt.call_id, elapsed_ms(attempt.started)) + .await + { + attempt + .first_valid_response_recorded + .store(false, Ordering::Release); + return Err(error); + } + } + drop(attempt); + match event { - ModelEvent::TextDelta(_) => { + ModelEvent::TextDelta(delta) if !delta.trim().is_empty() => { let attempt = self.inner.attempt.lock().await; if attempt .first_text_recorded @@ -470,6 +493,54 @@ mod tests { assert_eq!(count, 1); } + #[tokio::test] + async fn first_valid_response_includes_empty_text_and_reasoning_events() { + let store = Store::connect("sqlite::memory:").await.unwrap(); + let recorder = test_recorder(&store, "first-valid-response-call", false).await; + + recorder + .event(&ModelEvent::Start { + model_call_id: "call".into(), + }) + .await + .unwrap(); + recorder.event(&ModelEvent::TextStart).await.unwrap(); + recorder + .event(&ModelEvent::TextDelta(String::new())) + .await + .unwrap(); + + let call = store + .llm_call("first-valid-response-call") + .await + .unwrap() + .unwrap(); + assert!(call.ttfr_ms.is_some()); + assert!(call.ttft_ms.is_none()); + + recorder + .event(&ModelEvent::TextDelta("text".into())) + .await + .unwrap(); + let call = store + .llm_call("first-valid-response-call") + .await + .unwrap() + .unwrap(); + assert!(call.ttft_ms.is_some()); + + recorder + .event(&ModelEvent::ThinkingDelta("reasoning".into())) + .await + .unwrap(); + let call = store + .llm_call("first-valid-response-call") + .await + .unwrap() + .unwrap(); + assert!(call.first_valid_response_at_ms.is_some()); + } + #[tokio::test] async fn retry_finishes_the_old_call_and_records_the_new_request() { let store = Store::connect("sqlite::memory:").await.unwrap(); diff --git a/server/src/store/llm_calls.rs b/server/src/store/llm_calls.rs index 5f1ad5b..2511261 100644 --- a/server/src/store/llm_calls.rs +++ b/server/src/store/llm_calls.rs @@ -204,6 +204,21 @@ impl Store { Ok(()) } + pub async fn record_llm_first_valid_response( + &self, + call_id: &str, + elapsed_ms: i64, + ) -> Result<()> { + let _write = self.writes.lock().await; + sqlx::query("UPDATE llm_calls SET first_valid_response_at_ms = COALESCE(first_valid_response_at_ms, ?), ttfr_ms = COALESCE(ttfr_ms, ?) WHERE call_id = ?") + .bind(now_ms()) + .bind(elapsed_ms) + .bind(call_id) + .execute(&self.pool) + .await?; + Ok(()) + } + pub async fn record_llm_first_text(&self, call_id: &str, elapsed_ms: i64) -> Result<()> { let _write = self.writes.lock().await; sqlx::query("UPDATE llm_calls SET first_text_at_ms = COALESCE(first_text_at_ms, ?), ttft_ms = COALESCE(ttft_ms, ?) WHERE call_id = ?") @@ -369,10 +384,12 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result { response_headers_at_ms: row.try_get("response_headers_at_ms")?, first_event_at_ms: row.try_get("first_event_at_ms")?, first_text_at_ms: row.try_get("first_text_at_ms")?, + first_valid_response_at_ms: row.try_get("first_valid_response_at_ms")?, finished_at_ms: row.try_get("finished_at_ms")?, queue_ms: row.try_get("queue_ms")?, ttfb_ms: row.try_get("ttfb_ms")?, ttft_ms: row.try_get("ttft_ms")?, + ttfr_ms: row.try_get("ttfr_ms")?, duration_ms: row.try_get("duration_ms")?, input_tokens: row.try_get("input_tokens")?, output_tokens: row.try_get("output_tokens")?, diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 2e1a6b6..4051f35 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -207,7 +207,7 @@ async fn client_heartbeat_returns_a_server_protocol_heartbeat() { } #[tokio::test] -async fn runtime_user_message_action_aborts_active_exec_before_canceled_end_stream() { +async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() { let (_directory, store) = fixtures::temp_store().await; let provider = fake_provider::FakeProvider::default(); provider.push(vec![ @@ -281,7 +281,7 @@ async fn runtime_user_message_action_aborts_active_exec_before_canceled_end_stre handle .command(CursorCommand::Append { seqno: append_seqno, - message: Box::new(runtime_user_message()), + message: Box::new(runtime_cancel_action()), }) .await .unwrap(); @@ -313,6 +313,92 @@ async fn runtime_user_message_action_aborts_active_exec_before_canceled_end_stre assert_eq!(output.recv().await, None); } +#[tokio::test] +async fn runtime_user_message_action_interrupts_and_continues_with_new_message() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push_pending(); + provider.push(text_response("continued after user interruption")); + 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("user-message-request") + .await + .unwrap(); + let mut output = handle.subscribe(); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "user-message-request", + "user-message-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 + { + let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + if flags & connect::END_STREAM_FLAG != 0 { + panic!("initial run ended: {}", String::from_utf8_lossy(&payload)); + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + } + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_user_message()), + }) + .await + .unwrap(); + + let mut saw_continued = false; + let mut append_seqno = append_seqno + 1; + 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 { + if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message { + saw_continued |= delta.text.contains("continued after user interruption"); + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + assert!(saw_continued); + assert!(!handle.cancellation().is_cancelled()); + assert_eq!(provider.requests().len(), 2); + let history = serde_json::to_string(&provider.requests()[1].history).unwrap(); + assert!(history.contains("queued follow-up")); +} + #[tokio::test] async fn injected_user_context_restarts_only_the_active_model_cycle() { let (_directory, store) = fixtures::temp_store().await; @@ -1327,6 +1413,19 @@ fn kv_ack(id: u32) -> pb::AgentClientMessage { } } +fn runtime_cancel_action() -> pb::AgentClientMessage { + pb::AgentClientMessage { + message: Some(pb::agent_client_message::Message::ConversationAction( + pb::ConversationAction { + action: Some(pb::conversation_action::Action::CancelAction( + pb::CancelAction::default(), + )), + ..Default::default() + }, + )), + } +} + fn runtime_user_message() -> pb::AgentClientMessage { pb::AgentClientMessage { message: Some(pb::agent_client_message::Message::ConversationAction( diff --git a/server/tests/observability.rs b/server/tests/observability.rs index a2e8bb6..83f9b19 100644 --- a/server/tests/observability.rs +++ b/server/tests/observability.rs @@ -224,6 +224,7 @@ async fn records_one_summary_and_raw_payloads_for_one_provider_request() { assert_eq!(call.request_url, format!("http://{address}/proxy/generate")); assert_eq!(call.total_tokens, Some(12)); assert!(call.ttfb_ms.is_some()); + assert!(call.ttfr_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"); diff --git a/server/tests/schema_upgrade.rs b/server/tests/schema_upgrade.rs index 88ae3fe..a673303 100644 --- a/server/tests/schema_upgrade.rs +++ b/server/tests/schema_upgrade.rs @@ -40,6 +40,17 @@ async fn version_two_database_upgrades_with_cursor_request_mapping() { assert!(columns .iter() .any(|column| column.get::("name") == "cursor_request_id")); + + let llm_call_columns = sqlx::query("PRAGMA table_info(llm_calls)") + .fetch_all(store.pool()) + .await + .unwrap(); + assert!(llm_call_columns + .iter() + .any(|column| column.get::("name") == "first_valid_response_at_ms")); + assert!(llm_call_columns + .iter() + .any(|column| column.get::("name") == "ttfr_ms")); } #[tokio::test]