From df053c372079ef251c2668864c4d8c8836844fec Mon Sep 17 00:00:00 2001 From: leokun Date: Wed, 26 Aug 2026 16:46:41 +0800 Subject: [PATCH] fix: harden concurrent persistence and task recovery --- .github/workflows/release.yml | 20 ++ scripts/release/normalize-tauri-update.mjs | 97 +++++++++ server/src/cursor/observability.rs | 100 ++++++++- server/src/cursor/request/background.rs | 51 ++++- server/src/cursor/request/prepare.rs | 40 +++- server/src/provider/recorder.rs | 228 +++++++++++++++++++-- server/src/store/cas.rs | 21 +- server/src/store/cursor_traces.rs | 173 ++++++++++++++-- server/src/store/input_anchors.rs | 1 + server/src/store/llm_calls.rs | 225 ++++++++++++++++++-- server/src/store/mod.rs | 3 + server/src/store/models.rs | 5 + server/src/store/revisions.rs | 43 +++- server/src/store/runs.rs | 6 +- server/src/store/settings.rs | 17 +- server/src/store/sqlite.rs | 8 +- server/src/store/storage.rs | 1 + server/src/store/tool_rounds.rs | 5 +- server/src/store/writer.rs | 14 ++ server/tests/background_completion.rs | 73 ++++++- server/tests/observability.rs | 42 ++++ 21 files changed, 1083 insertions(+), 90 deletions(-) create mode 100644 scripts/release/normalize-tauri-update.mjs create mode 100644 server/src/store/writer.rs diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index f0bf863..682eff8 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -213,6 +213,26 @@ jobs: --output legacy-update/update.json \ --notes "Cursor BYOK v${VERSION}" + - name: Normalize Tauri updater download URLs + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + VERSION: ${{ needs.prepare.outputs.version }} + run: | + mkdir -p tauri-update + gh release download "v${VERSION}" --pattern latest.json --dir tauri-update + gh api "repos/${GITHUB_REPOSITORY}/releases/tags/v${VERSION}" > tauri-update/release.json + node scripts/release/normalize-tauri-update.mjs \ + --manifest tauri-update/latest.json \ + --release tauri-update/release.json \ + --repository "${GITHUB_REPOSITORY}" \ + --version "${VERSION}" + + - name: Upload normalized Tauri updater manifest + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + VERSION: ${{ needs.prepare.outputs.version }} + run: gh release upload "v${VERSION}" tauri-update/latest.json --clobber + - name: Upload legacy updater assets env: GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} diff --git a/scripts/release/normalize-tauri-update.mjs b/scripts/release/normalize-tauri-update.mjs new file mode 100644 index 0000000..9876237 --- /dev/null +++ b/scripts/release/normalize-tauri-update.mjs @@ -0,0 +1,97 @@ +import { readFile, writeFile } from "node:fs/promises"; +import { resolve } from "node:path"; +import { pathToFileURL } from "node:url"; + +function readOptions(args) { + const options = new Map(); + for (let index = 0; index < args.length; index += 2) { + const name = args[index]; + const value = args[index + 1]; + if (!name?.startsWith("--") || value === undefined) { + throw new Error(`invalid argument near ${name ?? "end of command"}`); + } + options.set(name.slice(2), value); + } + return options; +} + +function required(options, name) { + const value = options.get(name)?.trim(); + if (!value) throw new Error(`--${name} is required`); + return value; +} + +export function normalizeTauriUpdate(manifest, release, repository, version) { + if (manifest.version !== version) { + throw new Error( + `updater manifest version ${manifest.version ?? "is missing"}; expected ${version}`, + ); + } + if (release.tag_name !== `v${version}`) { + throw new Error( + `GitHub release tag ${release.tag_name ?? "is missing"}; expected v${version}`, + ); + } + if (!manifest.platforms || typeof manifest.platforms !== "object") { + throw new Error("updater manifest has no platforms"); + } + + const assetsByApiUrl = new Map(); + const publicAssetUrls = new Set(); + for (const asset of release.assets ?? []) { + if (!asset?.id || !asset?.browser_download_url) continue; + assetsByApiUrl.set( + `https://api.github.com/repos/${repository}/releases/assets/${asset.id}`, + asset.browser_download_url, + ); + publicAssetUrls.add(asset.browser_download_url); + } + + for (const [platform, entry] of Object.entries(manifest.platforms)) { + if (!entry?.signature || !entry?.url) { + throw new Error(`updater platform ${platform} is missing its URL or signature`); + } + const publicUrl = assetsByApiUrl.get(entry.url) ?? entry.url; + if (!publicAssetUrls.has(publicUrl)) { + throw new Error(`updater platform ${platform} references an unknown release asset`); + } + entry.url = publicUrl; + } + + return manifest; +} + +async function main() { + const options = readOptions(process.argv.slice(2)); + const manifestPath = resolve(required(options, "manifest")); + const releasePath = resolve(required(options, "release")); + const repository = required(options, "repository"); + const version = required(options, "version").replace(/^v/, ""); + + if (!/^[^/\s]+\/[^/\s]+$/.test(repository)) { + throw new Error(`invalid GitHub repository: ${repository}`); + } + if (!/^\d+\.\d+\.\d+(?:-[0-9A-Za-z.-]+)?$/.test(version)) { + throw new Error(`invalid semantic version: ${version}`); + } + + const manifest = JSON.parse(await readFile(manifestPath, "utf8")); + const release = JSON.parse(await readFile(releasePath, "utf8")); + const normalized = normalizeTauriUpdate( + manifest, + release, + repository, + version, + ); + await writeFile(manifestPath, `${JSON.stringify(normalized, null, 2)}\n`); +} + +if ( + process.argv[1] && + import.meta.url === pathToFileURL(resolve(process.argv[1])).href +) { + main().catch((error) => { + console.error(error instanceof Error ? error.message : String(error)); + process.exitCode = 1; + }); +} diff --git a/server/src/cursor/observability.rs b/server/src/cursor/observability.rs index f1244e7..d258944 100644 --- a/server/src/cursor/observability.rs +++ b/server/src/cursor/observability.rs @@ -1,11 +1,35 @@ -use crate::{store::BlobId, store::Store}; +use std::{ + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, + time::{Duration, Instant}, +}; + +use tokio::sync::Mutex; + +use crate::store::{BlobId, BufferedCursorTraceChunk, Store}; #[derive(Clone)] pub struct CursorTraceRecorder { store: Store, request_id: String, + chunks: Arc>, + finished: Arc, } +#[derive(Default)] +struct TraceChunkBuffer { + chunks: Vec, + bytes: usize, + first_chunk_at: Option, + generation: u64, +} + +const MAX_BUFFERED_CHUNKS: usize = 32; +const MAX_BUFFERED_BYTES: usize = 256 * 1024; +const MAX_BUFFER_AGE: Duration = Duration::from_millis(50); + impl CursorTraceRecorder { pub async fn begin( store: Store, @@ -21,6 +45,8 @@ impl CursorTraceRecorder { Ok(true) => Some(Self { store, request_id: request_id.into(), + chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())), + finished: Arc::new(AtomicBool::new(false)), }), Ok(false) => None, Err(error) => { @@ -35,6 +61,8 @@ impl CursorTraceRecorder { Ok(true) => Some(Self { store, request_id: request_id.into(), + chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())), + finished: Arc::new(AtomicBool::new(false)), }), Ok(false) => None, Err(error) => { @@ -115,16 +143,56 @@ impl CursorTraceRecorder { } pub async fn response_chunk(&self, source: &str, data: &[u8]) { - if let Err(error) = self - .store - .add_cursor_trace_response_chunk(&self.request_id, source, data) - .await + let mut buffer = self.chunks.lock().await; + if self.finished.load(Ordering::Acquire) { + return; + } + let schedule_flush = if buffer.chunks.is_empty() { + buffer.generation = buffer.generation.wrapping_add(1); + buffer.first_chunk_at = Some(Instant::now()); + Some(buffer.generation) + } else { + None + }; + buffer.bytes += data.len(); + buffer + .chunks + .push(BufferedCursorTraceChunk::new(source, data)); + let expired = buffer + .first_chunk_at + .is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE); + if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS + || buffer.bytes >= MAX_BUFFERED_BYTES + || expired { - tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk"); + if let Err(error) = self.flush_locked(&mut buffer).await { + tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk"); + } + } + drop(buffer); + if let Some(generation) = schedule_flush { + let recorder = self.clone(); + tokio::spawn(async move { + tokio::time::sleep(MAX_BUFFER_AGE).await; + let mut buffer = recorder.chunks.lock().await; + if buffer.generation == generation { + if let Err(error) = recorder.flush_locked(&mut buffer).await { + tracing::warn!(request_id = recorder.request_id, %error, "failed to flush Cursor response chunks"); + } + } + }); } } pub async fn finish(&self, error: Option<&str>) { + if self.finished.swap(true, Ordering::AcqRel) { + return; + } + let mut buffer = self.chunks.lock().await; + if let Err(store_error) = self.flush_locked(&mut buffer).await { + tracing::warn!(request_id = self.request_id, %store_error, "failed to flush Cursor response chunks"); + } + drop(buffer); if let Err(store_error) = self .store .finish_cursor_trace(&self.request_id, error) @@ -133,4 +201,24 @@ impl CursorTraceRecorder { tracing::warn!(request_id = self.request_id, %store_error, "failed to finish Cursor trace"); } } + + async fn flush_locked(&self, buffer: &mut TraceChunkBuffer) -> crate::Result<()> { + if buffer.chunks.is_empty() { + return Ok(()); + } + let chunks = std::mem::take(&mut buffer.chunks); + buffer.bytes = 0; + buffer.first_chunk_at = None; + if let Err(error) = self + .store + .add_cursor_trace_response_chunks(&self.request_id, &chunks) + .await + { + buffer.bytes = chunks.iter().map(|chunk| chunk.data.len()).sum(); + buffer.first_chunk_at = Some(Instant::now()); + buffer.chunks = chunks; + return Err(error); + } + Ok(()) + } } diff --git a/server/src/cursor/request/background.rs b/server/src/cursor/request/background.rs index eff4bed..08fd638 100644 --- a/server/src/cursor/request/background.rs +++ b/server/src/cursor/request/background.rs @@ -1,4 +1,4 @@ -use std::collections::BTreeSet; +use std::collections::BTreeMap; use crate::{cursor::proto::agent::v1 as pb, Error, Result}; @@ -35,8 +35,7 @@ pub(super) fn project( )); } - let mut identities = BTreeSet::new(); - let mut contexts = Vec::with_capacity(action.completions.len()); + let mut completions = BTreeMap::new(); let mut has_shell = false; let mut has_subagent = false; for completion in &action.completions { @@ -90,15 +89,21 @@ pub(super) fn project( }; let identity = agent_id.unwrap_or(&completion.task_id); let identity = format!("{}:{identity}", kind.as_str_name()); - if !identities.insert(identity.clone()) { + let context = completion_context(completion, kind, agent_id)?; + if completions + .insert(identity.clone(), (completion, context)) + .is_some() + { return Err(Error::Protocol(format!( "duplicate background task completion: {identity}" ))); } - contexts.push(completion_context(completion, kind, agent_id)?); } - let first = &action.completions[0]; + let (first, _) = completions + .values() + .next() + .expect("background completion action was validated as non-empty"); let text = match (has_shell, has_subagent) { (true, false) => SHELL_FOLLOW_UP.into(), (false, true) => FOLLOW_UP.into(), @@ -106,12 +111,16 @@ pub(super) fn project( (false, false) => unreachable!(), }; Ok(Projection { - context: contexts.join("\n\n"), + context: completions + .values() + .map(|(_, context)| context.as_str()) + .collect::>() + .join("\n\n"), turn_user: pb::UserMessage { text, message_id: format!( "background-completed:{}", - identities.into_iter().collect::>().join(":") + completions.keys().cloned().collect::>().join(":") ), mode, is_simulated_msg: Some(true), @@ -272,6 +281,32 @@ mod tests { assert!(projection.turn_user.text.contains(FOLLOW_UP)); } + #[test] + fn completion_batch_projection_is_independent_of_input_order() { + let first = completion(); + let mut second = completion(); + second.task_id = "child-id-2".into(); + second.subagent_id = Some("child-id-2".into()); + second.tool_call_id = Some("task-call-2".into()); + let forward = project( + &pb::BackgroundTaskCompletionAction { + completions: vec![first.clone(), second.clone()], + }, + pb::AgentMode::Multitask as i32, + ) + .unwrap(); + let reversed = project( + &pb::BackgroundTaskCompletionAction { + completions: vec![second, first], + }, + pb::AgentMode::Multitask as i32, + ) + .unwrap(); + + assert_eq!(forward.turn_user, reversed.turn_user); + assert_eq!(forward.context, reversed.context); + } + #[test] fn completion_requires_the_captured_subagent_identity_and_terminal_reason() { let mut value = completion(); diff --git a/server/src/cursor/request/prepare.rs b/server/src/cursor/request/prepare.rs index 00a27da..dd7883b 100644 --- a/server/src/cursor/request/prepare.rs +++ b/server/src/cursor/request/prepare.rs @@ -181,7 +181,7 @@ pub(crate) async fn prepare( None => proposed_base_revision_id, }; let existing_runtime = match event_id.as_deref() { - Some(event_id) if !background_completion => { + Some(event_id) => { store .message(&conversation_id, &format!("runtime:{event_id}")) .await? @@ -207,14 +207,22 @@ pub(crate) async fn prepare( } else { match (turn_user.clone(), event_id) { (Some(mut user), Some(event_id)) if background_completion => { - let (message, text) = runtime::compile_background( - event_id, - &user, - &request_context, - &action_context, - blob_sync, - ) - .await?; + let (message, text) = match existing_runtime { + Some(message) => { + let text = runtime_message_text(&message)?; + (message, text) + } + None => { + runtime::compile_background( + event_id, + &user, + &request_context, + &action_context, + blob_sync, + ) + .await? + } + }; user.text = text; turn_user = Some(user); vec![message] @@ -306,6 +314,20 @@ pub(crate) async fn prepare( )) } +fn runtime_message_text(message: &CanonicalMessage) -> Result { + let MessageContent::Parts { parts } = &message.content else { + return Err(Error::Protocol( + "stored runtime message does not contain parts".into(), + )); + }; + let Some(ContentPart::Text { text }) = parts.first() else { + return Err(Error::Protocol( + "stored runtime message does not start with text".into(), + )); + }; + Ok(text.clone()) +} + fn run_kind(subagent_type_name: Option<&str>, parent: Option<(RunId, String)>) -> Result { match (subagent_type_name, parent) { (None | Some("side-chat"), _) => Ok(RunKind::Root), diff --git a/server/src/provider/recorder.rs b/server/src/provider/recorder.rs index 005cbfe..1c80626 100644 --- a/server/src/provider/recorder.rs +++ b/server/src/provider/recorder.rs @@ -6,9 +6,11 @@ use std::{ time::Instant, }; +use tokio::sync::Mutex; + use crate::{ model::{NewLlmCall, Usage}, - store::Store, + store::{BufferedLlmChunk, Store}, Result, }; @@ -44,9 +46,23 @@ struct Inner { started: Instant, detailed: bool, next_chunk: AtomicI64, + chunks: Mutex, + first_text_recorded: AtomicBool, finished: AtomicBool, } +#[derive(Default)] +struct ChunkBuffer { + chunks: Vec, + bytes: usize, + first_chunk_at: Option, + generation: u64, +} + +const MAX_BUFFERED_CHUNKS: usize = 32; +const MAX_BUFFERED_BYTES: usize = 256 * 1024; +const MAX_BUFFER_AGE: std::time::Duration = std::time::Duration::from_millis(50); + impl CallRecorder { pub async fn start(store: Store, mut call: NewLlmCall) -> Result { call.detailed = store.detailed_logging().await?; @@ -58,6 +74,8 @@ impl CallRecorder { started: Instant::now(), detailed: call.detailed, next_chunk: AtomicI64::new(0), + chunks: Mutex::new(ChunkBuffer::default()), + first_text_recorded: AtomicBool::new(false), finished: AtomicBool::new(false), }), }) @@ -91,26 +109,67 @@ impl CallRecorder { } pub async fn response_chunk(&self, data: &[u8]) -> Result<()> { + let mut buffer = self.inner.chunks.lock().await; + if self.is_finished() { + return Ok(()); + } let seq = self.inner.next_chunk.fetch_add(1, Ordering::Relaxed); - self.inner - .store - .record_llm_chunk( - &self.inner.call_id, - seq, - self.elapsed_ms(), - data, - self.inner.detailed, - ) - .await + let schedule_flush = if buffer.chunks.is_empty() { + buffer.generation = buffer.generation.wrapping_add(1); + buffer.first_chunk_at = Some(Instant::now()); + Some(buffer.generation) + } else { + None + }; + buffer.bytes += data.len(); + buffer.chunks.push(if self.inner.detailed { + BufferedLlmChunk::new(seq, self.elapsed_ms(), data) + } else { + BufferedLlmChunk::metrics(seq, self.elapsed_ms(), data.len()) + }); + let expired = buffer + .first_chunk_at + .is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE); + if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS + || buffer.bytes >= MAX_BUFFERED_BYTES + || expired + { + self.flush_locked(&mut buffer).await?; + } + drop(buffer); + if let Some(generation) = schedule_flush { + let recorder = self.clone(); + tokio::spawn(async move { + tokio::time::sleep(MAX_BUFFER_AGE).await; + if let Err(error) = recorder.flush_generation(generation).await { + tracing::warn!(call_id = recorder.inner.call_id, %error, "failed to flush LLM response chunks"); + } + }); + } + Ok(()) } pub async fn event(&self, event: &ModelEvent) -> Result<()> { match event { ModelEvent::TextDelta(_) => { - self.inner - .store - .record_llm_first_text(&self.inner.call_id, self.elapsed_ms()) - .await?; + if self + .inner + .first_text_recorded + .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + if let Err(error) = self + .inner + .store + .record_llm_first_text(&self.inner.call_id, self.elapsed_ms()) + .await + { + self.inner + .first_text_recorded + .store(false, Ordering::Release); + return Err(error); + } + } } ModelEvent::Usage(usage) => self.usage(*usage).await?, ModelEvent::Done(reason) => self.completed(*reason).await?, @@ -155,6 +214,10 @@ impl CallRecorder { if self.inner.finished.swap(true, Ordering::AcqRel) { return Ok(()); } + if let Err(error) = self.flush_chunks().await { + self.inner.finished.store(false, Ordering::Release); + return Err(error); + } self.inner .store .finish_llm_call( @@ -168,6 +231,40 @@ impl CallRecorder { .await } + async fn flush_chunks(&self) -> Result<()> { + let mut buffer = self.inner.chunks.lock().await; + self.flush_locked(&mut buffer).await + } + + async fn flush_generation(&self, generation: u64) -> Result<()> { + let mut buffer = self.inner.chunks.lock().await; + if buffer.generation != generation { + return Ok(()); + } + self.flush_locked(&mut buffer).await + } + + async fn flush_locked(&self, buffer: &mut ChunkBuffer) -> Result<()> { + if buffer.chunks.is_empty() { + return Ok(()); + } + let chunks = std::mem::take(&mut buffer.chunks); + buffer.bytes = 0; + buffer.first_chunk_at = None; + if let Err(error) = self + .inner + .store + .record_llm_chunks(&self.inner.call_id, &chunks, self.inner.detailed) + .await + { + buffer.bytes = chunks.iter().map(|chunk| chunk.byte_count).sum(); + buffer.first_chunk_at = Some(Instant::now()); + buffer.chunks = chunks; + return Err(error); + } + Ok(()) + } + fn elapsed_ms(&self) -> i64 { self.inner .started @@ -193,3 +290,104 @@ fn error_kind(error: &crate::Error) -> &'static str { _ => "internal", } } + +#[cfg(test)] +mod tests { + use super::*; + + async fn test_recorder(store: &Store, call_id: &str, detailed: bool) -> CallRecorder { + sqlx::query( + "INSERT INTO llm_calls( + call_id, run_id, conversation_id, provider_call_index, provider_type, + provider_url, request_type, request_url, model_id, display_name, status, + created_at_ms, message_count, tool_count, detailed + ) VALUES (?, 'run', 'conversation', 0, 'openai-chat', + 'https://example.com', 'openai-chat', 'https://example.com', + 'model', 'Model', 'running', 1, 0, 0, ?)", + ) + .bind(call_id) + .bind(detailed) + .execute(store.pool()) + .await + .unwrap(); + CallRecorder { + inner: Arc::new(Inner { + store: store.clone(), + call_id: call_id.into(), + started: Instant::now(), + detailed, + next_chunk: AtomicI64::new(0), + chunks: Mutex::new(ChunkBuffer::default()), + first_text_recorded: AtomicBool::new(false), + finished: AtomicBool::new(false), + }), + } + } + + #[tokio::test] + async fn a_partial_chunk_batch_flushes_after_the_deadline() { + let store = Store::connect("sqlite::memory:").await.unwrap(); + let recorder = test_recorder(&store, "timed-flush-call", true).await; + + recorder.response_chunk(b"chunk").await.unwrap(); + assert_eq!( + store + .llm_call("timed-flush-call") + .await + .unwrap() + .unwrap() + .stream_event_count, + 0 + ); + + tokio::time::sleep(MAX_BUFFER_AGE + std::time::Duration::from_millis(100)).await; + + let call = store.llm_call("timed-flush-call").await.unwrap().unwrap(); + assert_eq!(call.response_bytes, 5); + assert_eq!(call.stream_event_count, 1); + assert_eq!( + store + .llm_call_chunks("timed-flush-call") + .await + .unwrap() + .len(), + 1 + ); + } + + #[tokio::test] + async fn first_text_is_persisted_only_once() { + let store = Store::connect("sqlite::memory:").await.unwrap(); + let recorder = test_recorder(&store, "first-text-call", false).await; + sqlx::query("CREATE TABLE first_text_updates(count INTEGER NOT NULL)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query("INSERT INTO first_text_updates(count) VALUES (0)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query( + "CREATE TRIGGER count_first_text_updates + AFTER UPDATE OF first_text_at_ms ON llm_calls + BEGIN + UPDATE first_text_updates SET count = count + 1; + END", + ) + .execute(store.pool()) + .await + .unwrap(); + for text in ["one", "two", "three"] { + recorder + .event(&ModelEvent::TextDelta(text.into())) + .await + .unwrap(); + } + + let count: i64 = sqlx::query_scalar("SELECT count FROM first_text_updates") + .fetch_one(store.pool()) + .await + .unwrap(); + assert_eq!(count, 1); + } +} diff --git a/server/src/store/cas.rs b/server/src/store/cas.rs index 447e521..f915178 100644 --- a/server/src/store/cas.rs +++ b/server/src/store/cas.rs @@ -1,6 +1,6 @@ use base64::{engine::general_purpose::STANDARD, Engine}; use sha2::{Digest, Sha256}; -use sqlx::Row; +use sqlx::{Row, Sqlite, Transaction}; use crate::{Error, Result}; @@ -44,13 +44,25 @@ pub struct BlobEdge { impl Store { pub async fn put_blob(&self, data: &[u8], edges: &[BlobEdge]) -> Result { + let _write = self.writes.lock().await; let blob_id = BlobId::digest(data); let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + Self::put_blob_tx(&mut tx, &blob_id, data, edges).await?; + tx.commit().await?; + Ok(blob_id) + } + + pub(crate) async fn put_blob_tx( + tx: &mut Transaction<'_, Sqlite>, + blob_id: &BlobId, + data: &[u8], + edges: &[BlobEdge], + ) -> Result<()> { sqlx::query("INSERT OR IGNORE INTO blobs(blob_id, data, created_at_ms) VALUES (?, ?, ?)") .bind(blob_id.as_bytes().as_slice()) .bind(data) .bind(now_ms()) - .execute(&mut *tx) + .execute(&mut **tx) .await?; for edge in edges { sqlx::query( @@ -59,11 +71,10 @@ impl Store { .bind(blob_id.as_bytes().as_slice()) .bind(edge.child.as_bytes().as_slice()) .bind(&edge.field_name) - .execute(&mut *tx) + .execute(&mut **tx) .await?; } - tx.commit().await?; - Ok(blob_id) + Ok(()) } pub async fn get_blob(&self, blob_id: &BlobId) -> Result>> { diff --git a/server/src/store/cursor_traces.rs b/server/src/store/cursor_traces.rs index d4f87ef..b4887a7 100644 --- a/server/src/store/cursor_traces.rs +++ b/server/src/store/cursor_traces.rs @@ -1,4 +1,4 @@ -use sqlx::Row; +use sqlx::{Row, Sqlite, Transaction}; use crate::{ model::{CursorRunTraceArtifact, CursorRunTraceSummary}, @@ -7,6 +7,21 @@ use crate::{ use super::{now_ms, BlobId, Store}; +#[derive(Clone, Debug)] +pub(crate) struct BufferedCursorTraceChunk { + pub(crate) source: String, + pub(crate) data: Vec, +} + +impl BufferedCursorTraceChunk { + pub(crate) fn new(source: &str, data: &[u8]) -> Self { + Self { + source: source.into(), + data: data.to_vec(), + } + } +} + impl Store { pub async fn start_cursor_trace_if_detailed( &self, @@ -21,6 +36,7 @@ impl Store { if !self.detailed_logging().await? { return Ok(false); } + let _write = self.writes.lock().await; sqlx::query( "INSERT OR IGNORE INTO cursor_run_traces( request_id, conversation_id, route, model_id, status, received_at_ms @@ -53,9 +69,22 @@ impl Store { data: &[u8], metadata: &serde_json::Value, ) -> Result<()> { - let blob_id = self.put_blob(data, &[]).await?; - self.link_cursor_trace_artifact(request_id, artifact_type, source, &blob_id, metadata) - .await + let metadata_json = serde_json::to_string(metadata)?; + let blob_id = BlobId::digest(data); + let _write = self.writes.lock().await; + let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + Self::put_blob_tx(&mut tx, &blob_id, data, &[]).await?; + Self::link_cursor_trace_artifact_tx( + &mut tx, + request_id, + artifact_type, + source, + &blob_id, + &metadata_json, + ) + .await?; + tx.commit().await?; + Ok(()) } pub async fn link_cursor_trace_artifact( @@ -66,13 +95,36 @@ impl Store { blob_id: &BlobId, metadata: &serde_json::Value, ) -> Result<()> { + let metadata_json = serde_json::to_string(metadata)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + Self::link_cursor_trace_artifact_tx( + &mut tx, + request_id, + artifact_type, + source, + blob_id, + &metadata_json, + ) + .await?; + tx.commit().await?; + Ok(()) + } + + async fn link_cursor_trace_artifact_tx( + tx: &mut Transaction<'_, Sqlite>, + request_id: &str, + artifact_type: &str, + source: &str, + blob_id: &BlobId, + metadata_json: &str, + ) -> Result<()> { let next: i64 = sqlx::query_scalar( "SELECT COALESCE(MAX(seq), -1) + 1 FROM cursor_run_trace_artifacts WHERE request_id = ?", ) .bind(request_id) - .fetch_one(&mut *tx) + .fetch_one(&mut **tx) .await?; sqlx::query( "INSERT INTO cursor_run_trace_artifacts( @@ -84,11 +136,10 @@ impl Store { .bind(artifact_type) .bind(source) .bind(blob_id.as_bytes().as_slice()) - .bind(serde_json::to_string(metadata)?) + .bind(metadata_json) .bind(now_ms()) - .execute(&mut *tx) + .execute(&mut **tx) .await?; - tx.commit().await?; Ok(()) } @@ -97,6 +148,7 @@ impl Store { request_id: &str, bytes: usize, ) -> Result<()> { + let _write = self.writes.lock().await; sqlx::query( "UPDATE cursor_run_traces SET request_bytes = request_bytes + ? WHERE request_id = ?", @@ -110,6 +162,7 @@ impl Store { pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> { let now = now_ms(); + let _write = self.writes.lock().await; sqlx::query( "UPDATE cursor_run_traces SET status = 'running', http_status = ?, @@ -131,30 +184,58 @@ impl Store { source: &str, data: &[u8], ) -> Result<()> { - self.append_cursor_trace_artifact( + self.add_cursor_trace_response_chunks( request_id, - "run_sse_chunk", - source, - data, - &serde_json::json!({"byte_count": data.len()}), + &[BufferedCursorTraceChunk::new(source, data)], ) - .await?; + .await + } + + pub(crate) async fn add_cursor_trace_response_chunks( + &self, + request_id: &str, + chunks: &[BufferedCursorTraceChunk], + ) -> Result<()> { + if chunks.is_empty() { + return Ok(()); + } + let response_bytes = chunks.iter().map(|chunk| chunk.data.len()).sum::(); + let _write = self.writes.lock().await; + let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + for chunk in chunks { + let metadata_json = + serde_json::to_string(&serde_json::json!({"byte_count": chunk.data.len()}))?; + let blob_id = BlobId::digest(&chunk.data); + Self::put_blob_tx(&mut tx, &blob_id, &chunk.data, &[]).await?; + Self::link_cursor_trace_artifact_tx( + &mut tx, + request_id, + "run_sse_chunk", + &chunk.source, + &blob_id, + &metadata_json, + ) + .await?; + } sqlx::query( "UPDATE cursor_run_traces SET response_bytes = response_bytes + ?, - response_event_count = response_event_count + 1, + response_event_count = response_event_count + ?, first_response_at_ms = COALESCE(first_response_at_ms, ?) WHERE request_id = ?", ) - .bind(as_i64(data.len())) + .bind(as_i64(response_bytes)) + .bind(chunks.len() as i64) .bind(now_ms()) .bind(request_id) - .execute(&self.pool) + .execute(&mut *tx) .await?; + tx.commit().await?; Ok(()) } pub async fn finish_cursor_trace(&self, request_id: &str, error: Option<&str>) -> Result<()> { + let _write = self.writes.lock().await; sqlx::query( "UPDATE cursor_run_traces SET status = ?, finished_at_ms = ?, error_message = ? @@ -244,3 +325,61 @@ fn trace_from_row(row: sqlx::sqlite::SqliteRow) -> Result fn as_i64(value: usize) -> i64 { value.min(i64::MAX as usize) as i64 } + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn records_a_batch_of_trace_chunks_with_one_summary_update() { + let store = Store::connect("sqlite::memory:").await.unwrap(); + store.set_detailed_logging(true).await.unwrap(); + store + .start_cursor_trace_if_detailed("trace", None, "cursor_official", None) + .await + .unwrap(); + sqlx::query("CREATE TABLE trace_summary_updates(count INTEGER NOT NULL)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query("INSERT INTO trace_summary_updates(count) VALUES (0)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query( + "CREATE TRIGGER count_trace_summary_updates + AFTER UPDATE OF response_bytes ON cursor_run_traces + BEGIN + UPDATE trace_summary_updates SET count = count + 1; + END", + ) + .execute(store.pool()) + .await + .unwrap(); + + store + .add_cursor_trace_response_chunks( + "trace", + &[ + BufferedCursorTraceChunk::new("cursor_official", b"one"), + BufferedCursorTraceChunk::new("cursor_official", b"two"), + BufferedCursorTraceChunk::new("cursor_official", b"three"), + ], + ) + .await + .unwrap(); + + let trace = store.cursor_trace("trace").await.unwrap().unwrap(); + assert_eq!(trace.response_bytes, 11); + assert_eq!(trace.response_event_count, 3); + assert_eq!( + store.cursor_trace_artifacts("trace").await.unwrap().len(), + 3 + ); + let updates: i64 = sqlx::query_scalar("SELECT count FROM trace_summary_updates") + .fetch_one(store.pool()) + .await + .unwrap(); + assert_eq!(updates, 1); + } +} diff --git a/server/src/store/input_anchors.rs b/server/src/store/input_anchors.rs index beb6f6d..d657726 100644 --- a/server/src/store/input_anchors.rs +++ b/server/src/store/input_anchors.rs @@ -12,6 +12,7 @@ impl Store { input_id: &str, base_revision_id: RevisionId, ) -> Result { + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; sqlx::query( "INSERT INTO input_anchors diff --git a/server/src/store/llm_calls.rs b/server/src/store/llm_calls.rs index db3ca2e..5f1ad5b 100644 --- a/server/src/store/llm_calls.rs +++ b/server/src/store/llm_calls.rs @@ -12,6 +12,34 @@ use crate::{ use super::{now_ms, Store}; +#[derive(Clone, Debug)] +pub(crate) struct BufferedLlmChunk { + pub(crate) seq: i64, + pub(crate) elapsed_ms: i64, + pub(crate) data: Option>, + pub(crate) byte_count: usize, +} + +impl BufferedLlmChunk { + pub(crate) fn new(seq: i64, elapsed_ms: i64, data: &[u8]) -> Self { + Self { + seq, + elapsed_ms, + data: Some(data.to_vec()), + byte_count: data.len(), + } + } + + pub(crate) fn metrics(seq: i64, elapsed_ms: i64, byte_count: usize) -> Self { + Self { + seq, + elapsed_ms, + data: None, + byte_count, + } + } +} + impl Store { pub async fn detailed_logging(&self) -> Result { let value: String = sqlx::query_scalar( @@ -23,6 +51,7 @@ impl Store { } pub async fn set_detailed_logging(&self, enabled: bool) -> Result<()> { + let _write = self.writes.lock().await; sqlx::query( "INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES ('llm_detailed_logging', ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms", ) @@ -34,6 +63,7 @@ impl Store { } pub async fn start_llm_call(&self, call: &NewLlmCall) -> Result<()> { + let _write = self.writes.lock().await; let now = now_ms(); sqlx::query( r#"INSERT INTO llm_calls( @@ -74,21 +104,27 @@ impl Store { detailed: bool, ) -> Result<()> { let body_json = serde_json::to_string(body)?; + let headers_json = detailed + .then(|| serde_json::to_string(headers)) + .transpose()?; + let _write = self.writes.lock().await; + let mut transaction = self.pool.begin_with("BEGIN IMMEDIATE").await?; if detailed { sqlx::query("INSERT INTO llm_call_requests(call_id, headers_json, body_json, byte_count) SELECT ?, ?, ?, ? WHERE EXISTS (SELECT 1 FROM llm_calls WHERE call_id = ?)") .bind(call_id) - .bind(serde_json::to_string(headers)?) + .bind(headers_json) .bind(&body_json) .bind(body_json.len() as i64) .bind(call_id) - .execute(&self.pool) + .execute(&mut *transaction) .await?; } sqlx::query("UPDATE llm_calls SET request_bytes = ? WHERE call_id = ?") .bind(body_json.len() as i64) .bind(call_id) - .execute(&self.pool) + .execute(&mut *transaction) .await?; + transaction.commit().await?; Ok(()) } @@ -98,6 +134,7 @@ impl Store { elapsed_ms: i64, http_status: u16, ) -> Result<()> { + let _write = self.writes.lock().await; sqlx::query("UPDATE llm_calls SET response_headers_at_ms = ?, ttfb_ms = ?, http_status = ? WHERE call_id = ?") .bind(now_ms()) .bind(elapsed_ms) @@ -116,21 +153,50 @@ impl Store { data: &[u8], detailed: bool, ) -> Result<()> { - let mut transaction = self.pool.begin().await?; - if detailed { - sqlx::query("INSERT INTO llm_call_response_chunks(call_id, seq, received_offset_ms, data, byte_count) SELECT ?, ?, ?, ?, ? WHERE EXISTS (SELECT 1 FROM llm_calls WHERE call_id = ?)") - .bind(call_id) - .bind(seq) - .bind(elapsed_ms) - .bind(data) - .bind(data.len() as i64) - .bind(call_id) - .execute(&mut *transaction) - .await?; + let chunk = if detailed { + BufferedLlmChunk::new(seq, elapsed_ms, data) + } else { + BufferedLlmChunk::metrics(seq, elapsed_ms, data.len()) + }; + self.record_llm_chunks(call_id, &[chunk], detailed).await + } + + pub(crate) async fn record_llm_chunks( + &self, + call_id: &str, + chunks: &[BufferedLlmChunk], + detailed: bool, + ) -> Result<()> { + if chunks.is_empty() { + return Ok(()); } - sqlx::query("UPDATE llm_calls SET first_event_at_ms = COALESCE(first_event_at_ms, ?), response_bytes = response_bytes + ?, stream_event_count = stream_event_count + 1 WHERE call_id = ?") + let byte_count = chunks + .iter() + .map(|chunk| chunk.byte_count as i64) + .sum::(); + let event_count = chunks.len() as i64; + let _write = self.writes.lock().await; + let mut transaction = self.pool.begin_with("BEGIN IMMEDIATE").await?; + if detailed { + for chunk in chunks { + let data = chunk.data.as_deref().ok_or_else(|| { + crate::Error::Store("detailed LLM chunk is missing payload data".into()) + })?; + sqlx::query("INSERT INTO llm_call_response_chunks(call_id, seq, received_offset_ms, data, byte_count) SELECT ?, ?, ?, ?, ? WHERE EXISTS (SELECT 1 FROM llm_calls WHERE call_id = ?)") + .bind(call_id) + .bind(chunk.seq) + .bind(chunk.elapsed_ms) + .bind(data) + .bind(chunk.byte_count as i64) + .bind(call_id) + .execute(&mut *transaction) + .await?; + } + } + sqlx::query("UPDATE llm_calls SET first_event_at_ms = COALESCE(first_event_at_ms, ?), response_bytes = response_bytes + ?, stream_event_count = stream_event_count + ? WHERE call_id = ?") .bind(now_ms()) - .bind(data.len() as i64) + .bind(byte_count) + .bind(event_count) .bind(call_id) .execute(&mut *transaction) .await?; @@ -139,6 +205,7 @@ impl Store { } 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 = ?") .bind(now_ms()) .bind(elapsed_ms) @@ -149,6 +216,8 @@ impl Store { } pub async fn record_llm_usage(&self, call_id: &str, usage: Usage) -> Result<()> { + let usage_json = serde_json::to_string(&usage)?; + let _write = self.writes.lock().await; sqlx::query("UPDATE llm_calls SET input_tokens = ?, output_tokens = ?, total_tokens = ?, cache_read_tokens = ?, cache_write_tokens = ?, reasoning_tokens = ?, usage_json = ? WHERE call_id = ?") .bind(as_i64(usage.input_tokens)) .bind(as_i64(usage.output_tokens)) @@ -156,7 +225,7 @@ impl Store { .bind(as_i64(usage.cache_read_tokens)) .bind(as_i64(usage.cache_write_tokens)) .bind(as_i64(usage.reasoning_tokens)) - .bind(serde_json::to_string(&usage)?) + .bind(usage_json) .bind(call_id) .execute(&self.pool) .await?; @@ -172,6 +241,7 @@ impl Store { error_kind: Option<&str>, error_message: Option<&str>, ) -> Result<()> { + let _write = self.writes.lock().await; sqlx::query("UPDATE llm_calls SET status = ?, finish_reason = ?, finished_at_ms = ?, duration_ms = ?, error_kind = ?, error_message = ? WHERE call_id = ? AND status = 'running'") .bind(status) .bind(finish_reason) @@ -327,8 +397,129 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result { #[cfg(test)] mod tests { + use std::sync::Arc; + use super::*; use crate::model::{ModelConfigInput, ModelType}; + use tokio::sync::Barrier; + + #[tokio::test] + async fn concurrent_writes_are_serialized_without_sqlite_busy_retries() { + let directory = tempfile::tempdir().unwrap(); + let store = Store::connect(&format!( + "sqlite://{}", + directory.path().join("concurrent-writes.db").display() + )) + .await + .unwrap(); + sqlx::query( + "INSERT INTO llm_calls( + call_id, run_id, conversation_id, provider_call_index, provider_type, + provider_url, request_type, request_url, model_id, display_name, status, + created_at_ms, message_count, tool_count, detailed + ) VALUES ( + 'concurrent-call', 'run', 'conversation', 0, 'openai-chat', + 'https://example.com', 'openai-chat', 'https://example.com', + 'model', 'Model', 'running', 1, 0, 0, 0 + )", + ) + .execute(store.pool()) + .await + .unwrap(); + + let mut connections = Vec::new(); + for _ in 0..8 { + connections.push(store.pool().acquire().await.unwrap()); + } + for connection in &mut connections { + sqlx::query("PRAGMA busy_timeout = 0") + .execute(&mut **connection) + .await + .unwrap(); + } + drop(connections); + + let writers = 32; + let barrier = Arc::new(Barrier::new(writers)); + let mut tasks = Vec::with_capacity(writers); + for seq in 0..writers { + let store = store.clone(); + let barrier = barrier.clone(); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + store + .record_llm_chunk("concurrent-call", seq as i64, 1, b"x", false) + .await + })); + } + for task in tasks { + task.await.unwrap().unwrap(); + } + + let call = store.llm_call("concurrent-call").await.unwrap().unwrap(); + assert_eq!(call.response_bytes, writers as i64); + assert_eq!(call.stream_event_count, writers as i64); + } + + #[tokio::test] + async fn records_a_batch_of_response_chunks_with_one_summary_update() { + let store = Store::connect("sqlite::memory:").await.unwrap(); + sqlx::query( + "INSERT INTO llm_calls( + call_id, run_id, conversation_id, provider_call_index, provider_type, + provider_url, request_type, request_url, model_id, display_name, status, + created_at_ms, message_count, tool_count, detailed + ) VALUES ( + 'batch-call', 'run', 'conversation', 0, 'openai-chat', + 'https://example.com', 'openai-chat', 'https://example.com', + 'model', 'Model', 'running', 1, 0, 0, 1 + )", + ) + .execute(store.pool()) + .await + .unwrap(); + sqlx::query("CREATE TABLE llm_call_summary_updates(count INTEGER NOT NULL)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query("INSERT INTO llm_call_summary_updates(count) VALUES (0)") + .execute(store.pool()) + .await + .unwrap(); + sqlx::query( + "CREATE TRIGGER count_llm_call_summary_updates + AFTER UPDATE OF response_bytes ON llm_calls + BEGIN + UPDATE llm_call_summary_updates SET count = count + 1; + END", + ) + .execute(store.pool()) + .await + .unwrap(); + + store + .record_llm_chunks( + "batch-call", + &[ + BufferedLlmChunk::new(0, 1, b"one"), + BufferedLlmChunk::new(1, 2, b"two"), + BufferedLlmChunk::new(2, 3, b"three"), + ], + true, + ) + .await + .unwrap(); + + let call = store.llm_call("batch-call").await.unwrap().unwrap(); + assert_eq!(call.response_bytes, 11); + assert_eq!(call.stream_event_count, 3); + assert_eq!(store.llm_call_chunks("batch-call").await.unwrap().len(), 3); + let updates: i64 = sqlx::query_scalar("SELECT count FROM llm_call_summary_updates") + .fetch_one(store.pool()) + .await + .unwrap(); + assert_eq!(updates, 1); + } #[tokio::test] async fn latest_usage_anchor_uses_the_latest_completed_call_for_the_same_conversation_and_model( diff --git a/server/src/store/mod.rs b/server/src/store/mod.rs index 9a21a16..85d7739 100644 --- a/server/src/store/mod.rs +++ b/server/src/store/mod.rs @@ -13,8 +13,11 @@ mod settings; mod sqlite; mod storage; mod tool_rounds; +mod writer; pub use cas::*; +pub(crate) use cursor_traces::BufferedCursorTraceChunk; +pub(crate) use llm_calls::BufferedLlmChunk; pub use runs::*; pub use settings::*; pub(crate) use sqlite::now_ms; diff --git a/server/src/store/models.rs b/server/src/store/models.rs index d26b983..2a7351d 100644 --- a/server/src/store/models.rs +++ b/server/src/store/models.rs @@ -60,6 +60,7 @@ impl Store { normalized.push((hash, input)); } let now = now_ms(); + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; for (hash, input) in &normalized { insert_model(&mut transaction, hash, input, now).await?; @@ -87,6 +88,7 @@ impl Store { } } let now = now_ms(); + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; let mut inserted = 0; for (hash, input) in &normalized { @@ -110,6 +112,7 @@ impl Store { let input = normalize_model_input(input)?; let next_hash = model_hash(&input)?; let now = now_ms(); + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; if next_hash != current.model_hash { sqlx::query("UPDATE llm_calls SET model_hash = NULL WHERE model_hash = ?") @@ -165,6 +168,7 @@ impl Store { } pub async fn delete_model(&self, hash: &str) -> Result<()> { + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; sqlx::query("UPDATE llm_calls SET model_hash = NULL WHERE model_hash = ?") .bind(hash) @@ -201,6 +205,7 @@ impl Store { } let now = now_ms(); + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; for (index, hash) in model_hashes.iter().enumerate() { sqlx::query( diff --git a/server/src/store/revisions.rs b/server/src/store/revisions.rs index 8bc1a7b..5e638fd 100644 --- a/server/src/store/revisions.rs +++ b/server/src/store/revisions.rs @@ -13,6 +13,7 @@ impl Store { &self, conversation_id: &ConversationId, ) -> Result { + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let revision = Self::ensure_conversation_tx(&mut tx, conversation_id).await?; tx.commit().await?; @@ -62,9 +63,10 @@ impl Store { conversation_id: &ConversationId, messages: &[CanonicalMessage], ) -> Result { + let digest = message_digest(messages)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let current = Self::ensure_conversation_tx(&mut tx, conversation_id).await?; - let digest = message_digest(messages)?; if let Some(existing) = sqlx::query_scalar::<_, i64>( "SELECT revision_id FROM conversation_revisions WHERE conversation_id = ? AND state_digest = ?", @@ -107,9 +109,20 @@ impl Store { if additions.is_empty() { return Ok(expected); } + let mut full = self.load_revision_messages(expected).await?; + full.extend_from_slice(additions); + let digest = message_digest(&full)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; - let revision = - Self::append_revision_tx(&mut tx, conversation_id, run_id, expected, additions).await?; + let revision = Self::append_revision_with_digest_tx( + &mut tx, + conversation_id, + run_id, + expected, + additions, + digest, + ) + .await?; tx.commit().await?; Ok(revision) } @@ -121,6 +134,8 @@ impl Store { expected: RevisionId, messages: &[CanonicalMessage], ) -> Result { + let digest = message_digest(messages)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; Self::require_active_head_tx(&mut tx, conversation_id, run_id, expected).await?; let root: i64 = sqlx::query_scalar( @@ -130,7 +145,6 @@ impl Store { .bind(conversation_id.as_str()) .fetch_one(&mut *tx) .await?; - let digest = message_digest(messages)?; let revision = Self::insert_revision_tx(&mut tx, conversation_id, RevisionId(root), messages, digest) .await?; @@ -202,10 +216,29 @@ impl Store { expected: RevisionId, additions: &[CanonicalMessage], ) -> Result { - Self::require_active_head_tx(tx, conversation_id, run_id, expected).await?; let mut full = Self::load_revision_messages_tx(tx, expected.0).await?; full.extend_from_slice(additions); let digest = message_digest(&full)?; + Self::append_revision_with_digest_tx( + tx, + conversation_id, + run_id, + expected, + additions, + digest, + ) + .await + } + + async fn append_revision_with_digest_tx( + tx: &mut Transaction<'_, Sqlite>, + conversation_id: &ConversationId, + run_id: &RunId, + expected: RevisionId, + additions: &[CanonicalMessage], + digest: [u8; 32], + ) -> Result { + Self::require_active_head_tx(tx, conversation_id, run_id, expected).await?; if sqlx::query_scalar::<_, i64>( "SELECT revision_id FROM conversation_revisions WHERE conversation_id = ? AND state_digest = ?", diff --git a/server/src/store/runs.rs b/server/src/store/runs.rs index 9c8a70e..d908b25 100644 --- a/server/src/store/runs.rs +++ b/server/src/store/runs.rs @@ -36,6 +36,7 @@ pub struct ClaimedRun { impl Store { pub async fn claim_run(&self, prepared: &PreparedRun) -> Result { + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let now = now_ms(); Self::ensure_conversation_tx(&mut tx, &prepared.conversation_id).await?; @@ -146,6 +147,7 @@ impl Store { } pub async fn begin_provider_call(&self, run_id: &RunId) -> Result { + let _write = self.writes.lock().await; let index: Option = sqlx::query_scalar( "UPDATE runs SET provider_call_index = provider_call_index + 1, updated_at_ms = ? WHERE run_id = ? AND status = 'running' @@ -167,6 +169,8 @@ impl Store { usage: Option, failure: Option<(&str, &str)>, ) -> Result { + let usage_json = serde_json::to_string(&usage)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let row = sqlx::query( "SELECT conversation_id, status, failure_category, failure_summary @@ -200,7 +204,7 @@ impl Store { WHERE run_id = ? AND status = 'running'", ) .bind(status.as_str()) - .bind(serde_json::to_string(&usage)?) + .bind(usage_json) .bind(category) .bind(summary) .bind(now) diff --git a/server/src/store/settings.rs b/server/src/store/settings.rs index ca8fbce..3377a46 100644 --- a/server/src/store/settings.rs +++ b/server/src/store/settings.rs @@ -93,6 +93,7 @@ pub(crate) struct ProxySettingsSecret { impl Store { pub(crate) async fn installation_id(&self) -> Result { let generated = uuid::Uuid::new_v4().to_string(); + let _write = self.writes.lock().await; sqlx::query( "INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO NOTHING", ) @@ -165,9 +166,11 @@ impl Store { username: input.username.trim().to_owned(), password, }; + let value_json = serde_json::to_string(&settings)?; + let _write = self.writes.lock().await; sqlx::query("INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms") .bind(PROXY_SETTINGS_KEY) - .bind(serde_json::to_string(&settings)?) + .bind(value_json) .bind(now_ms()) .execute(&self.pool) .await?; @@ -206,9 +209,11 @@ impl Store { )); } } + let value_json = serde_json::to_string(&settings)?; + let _write = self.writes.lock().await; sqlx::query("INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms") .bind(TAB_SETTINGS_KEY) - .bind(serde_json::to_string(&settings)?) + .bind(value_json) .bind(now_ms()) .execute(&self.pool) .await?; @@ -228,11 +233,13 @@ impl Store { } pub async fn set_port_settings(&self, settings: PortSettings) -> Result<()> { + let value_json = serde_json::to_string(&settings)?; + let _write = self.writes.lock().await; sqlx::query( "INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms", ) .bind(PORT_SETTINGS_KEY) - .bind(serde_json::to_string(&settings)?) + .bind(value_json) .bind(now_ms()) .execute(&self.pool) .await?; @@ -264,11 +271,13 @@ impl Store { } pub async fn set_desktop_settings(&self, settings: DesktopSettings) -> Result<()> { + let value_json = serde_json::to_string(&settings)?; + let _write = self.writes.lock().await; sqlx::query( "INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms", ) .bind(DESKTOP_SETTINGS_KEY) - .bind(serde_json::to_string(&settings)?) + .bind(value_json) .bind(now_ms()) .execute(&self.pool) .await?; diff --git a/server/src/store/sqlite.rs b/server/src/store/sqlite.rs index 8438108..06a2ec5 100644 --- a/server/src/store/sqlite.rs +++ b/server/src/store/sqlite.rs @@ -7,9 +7,12 @@ use sqlx::{ use crate::Result; +use super::writer::WriteCoordinator; + #[derive(Clone)] pub struct Store { pub(crate) pool: SqlitePool, + pub(crate) writes: WriteCoordinator, } impl Store { @@ -25,7 +28,10 @@ impl Store { .connect_with(options) .await?; sqlx::migrate!("./migrations").run(&pool).await?; - Ok(Self { pool }) + Ok(Self { + pool, + writes: WriteCoordinator::default(), + }) } pub fn pool(&self) -> &SqlitePool { diff --git a/server/src/store/storage.rs b/server/src/store/storage.rs index f5b1255..ee64bb6 100644 --- a/server/src/store/storage.rs +++ b/server/src/store/storage.rs @@ -54,6 +54,7 @@ impl Store { } pub async fn clear_statistics_storage(&self) -> Result { + let _write = self.writes.lock().await; let mut transaction = self.pool.begin().await?; sqlx::query("DELETE FROM llm_calls") .execute(&mut *transaction) diff --git a/server/src/store/tool_rounds.rs b/server/src/store/tool_rounds.rs index ba71d1a..95bb9c4 100644 --- a/server/src/store/tool_rounds.rs +++ b/server/src/store/tool_rounds.rs @@ -50,6 +50,8 @@ impl Store { if calls.is_empty() { return Err(Error::Store("cannot persist an empty tool round".into())); } + let assistant_json = serde_json::to_string(assistant)?; + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let ownership: bool = sqlx::query_scalar( "SELECT EXISTS( @@ -82,7 +84,7 @@ impl Store { .bind(round_id.as_str()) .bind(run_id.as_str()) .bind(base_revision_id.0) - .bind(serde_json::to_string(assistant)?) + .bind(assistant_json) .bind(created_at_ms) .bind(now) .execute(&mut *tx) @@ -113,6 +115,7 @@ impl Store { round_id: &ToolRoundId, result: &ToolResult, ) -> Result { + let _write = self.writes.lock().await; let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; let round = sqlx::query( "SELECT assistant_json, status, version, next_completion_seq diff --git a/server/src/store/writer.rs b/server/src/store/writer.rs new file mode 100644 index 0000000..b14afef --- /dev/null +++ b/server/src/store/writer.rs @@ -0,0 +1,14 @@ +use std::sync::Arc; + +use tokio::sync::{Mutex, MutexGuard}; + +#[derive(Clone, Default)] +pub(crate) struct WriteCoordinator { + lock: Arc>, +} + +impl WriteCoordinator { + pub(crate) async fn lock(&self) -> MutexGuard<'_, ()> { + self.lock.lock().await + } +} diff --git a/server/tests/background_completion.rs b/server/tests/background_completion.rs index 07adc15..b6e7efb 100644 --- a/server/tests/background_completion.rs +++ b/server/tests/background_completion.rs @@ -137,6 +137,68 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() { ); } +#[tokio::test] +async fn retrying_one_background_completion_reuses_its_runtime_message() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push(stop_response("model-call", "followed up")); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = CursorSessionRegistry::new( + store.clone(), + Arc::new(provider.clone()), + PromptCompiler::new(assets), + Default::default(), + ); + let first = registry.get_or_create("completion-retry-1").await.unwrap(); + let (checkpoint, _) = drive_completion( + &first, + completion_run( + "retry-child", + "completion-retry-run-1", + pb::ConversationStateStructure { + mode: Some(pb::AgentMode::Multitask as i32), + ..Default::default() + }, + ), + ) + .await; + + provider.push(stop_response("model-call-2", "followed up again")); + let second = registry.get_or_create("completion-retry-2").await.unwrap(); + drive_completion( + &second, + completion_run_with_detail( + "retry-child", + "completion-retry-run-2", + checkpoint, + "updated retry payload", + ), + ) + .await; + + let messages = store + .load_current_messages(&cursor_server::model::ConversationId::new( + "parent-conversation", + )) + .await + .unwrap(); + assert_eq!( + messages + .iter() + .filter(|message| { + message.runtime_event_id.as_deref() + == Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child") + }) + .count(), + 1 + ); +} + #[tokio::test] async fn background_shell_completion_wakes_the_parent_with_the_captured_notification() { let (_directory, store) = fixtures::temp_store().await; @@ -339,6 +401,15 @@ fn completion_run( child_id: &str, run_id: &str, conversation_state: pb::ConversationStateStructure, +) -> pb::AgentClientMessage { + completion_run_with_detail(child_id, run_id, conversation_state, "child result") +} + +fn completion_run_with_detail( + child_id: &str, + run_id: &str, + conversation_state: pb::ConversationStateStructure, + detail: &str, ) -> pb::AgentClientMessage { pb::AgentClientMessage { message: Some(pb::agent_client_message::Message::RunRequest( @@ -352,7 +423,7 @@ fn completion_run( kind: pb::BackgroundTaskKind::Subagent as i32, status: pb::BackgroundTaskStatus::Success as i32, title: "Inspect protocol".into(), - detail: Some("child result".into()), + detail: Some(detail.into()), output_path: Some("/tmp/child.jsonl".into()), reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32, subagent_id: Some(child_id.into()), diff --git a/server/tests/observability.rs b/server/tests/observability.rs index 6345cf7..a2e8bb6 100644 --- a/server/tests/observability.rs +++ b/server/tests/observability.rs @@ -102,6 +102,48 @@ async fn cursor_trace_links_detailed_artifacts_to_the_logical_run() { assert_eq!(artifacts[1].data, b"response"); } +#[tokio::test] +async fn cursor_trace_artifact_and_blob_are_written_atomically() { + let (_directory, store) = test_store("cursor-trace-atomic.db").await; + store.set_detailed_logging(true).await.unwrap(); + store + .start_cursor_trace_if_detailed( + "request-atomic", + Some("conversation"), + "cursor_official", + Some("model"), + ) + .await + .unwrap(); + sqlx::query( + "CREATE TRIGGER reject_trace_artifact + BEFORE INSERT ON cursor_run_trace_artifacts + BEGIN + SELECT RAISE(ABORT, 'rejected artifact'); + END", + ) + .execute(store.pool()) + .await + .unwrap(); + + assert!(store + .append_cursor_trace_artifact( + "request-atomic", + "run_sse_chunk", + "cursor_official", + b"must-rollback", + &serde_json::json!({}), + ) + .await + .is_err()); + + let blob_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM blobs") + .fetch_one(store.pool()) + .await + .unwrap(); + assert_eq!(blob_count, 0); +} + #[tokio::test] async fn records_one_summary_and_raw_payloads_for_one_provider_request() { let app = Router::new().route(