fix: harden concurrent persistence and task recovery

This commit is contained in:
leokun
2026-08-26 16:46:41 +08:00
parent 9deb42915c
commit df053c3720
21 changed files with 1083 additions and 90 deletions
+20
View File
@@ -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 }}
@@ -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;
});
}
+94 -6
View File
@@ -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<Mutex<TraceChunkBuffer>>,
finished: Arc<AtomicBool>,
}
#[derive(Default)]
struct TraceChunkBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
first_chunk_at: Option<Instant>,
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(())
}
}
+43 -8
View File
@@ -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::<Vec<_>>()
.join("\n\n"),
turn_user: pb::UserMessage {
text,
message_id: format!(
"background-completed:{}",
identities.into_iter().collect::<Vec<_>>().join(":")
completions.keys().cloned().collect::<Vec<_>>().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();
+31 -9
View File
@@ -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<String> {
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<RunKind> {
match (subagent_type_name, parent) {
(None | Some("side-chat"), _) => Ok(RunKind::Root),
+213 -15
View File
@@ -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<ChunkBuffer>,
first_text_recorded: AtomicBool,
finished: AtomicBool,
}
#[derive(Default)]
struct ChunkBuffer {
chunks: Vec<BufferedLlmChunk>,
bytes: usize,
first_chunk_at: Option<Instant>,
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<Self> {
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);
}
}
+16 -5
View File
@@ -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<BlobId> {
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<Option<Vec<u8>>> {
+156 -17
View File
@@ -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<u8>,
}
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::<usize>();
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<CursorRunTraceSummary>
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);
}
}
+1
View File
@@ -12,6 +12,7 @@ impl Store {
input_id: &str,
base_revision_id: RevisionId,
) -> Result<RevisionId> {
let _write = self.writes.lock().await;
let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?;
sqlx::query(
"INSERT INTO input_anchors
+208 -17
View File
@@ -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<Vec<u8>>,
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<bool> {
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::<i64>();
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<LlmCallSummary> {
#[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(
+3
View File
@@ -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;
+5
View File
@@ -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(
+38 -5
View File
@@ -13,6 +13,7 @@ impl Store {
&self,
conversation_id: &ConversationId,
) -> Result<RevisionId> {
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<RevisionId> {
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<RevisionId> {
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<RevisionId> {
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<RevisionId> {
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 = ?",
+5 -1
View File
@@ -36,6 +36,7 @@ pub struct ClaimedRun {
impl Store {
pub async fn claim_run(&self, prepared: &PreparedRun) -> Result<ClaimedRun> {
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<u64> {
let _write = self.writes.lock().await;
let index: Option<i64> = 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<Usage>,
failure: Option<(&str, &str)>,
) -> Result<bool> {
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)
+13 -4
View File
@@ -93,6 +93,7 @@ pub(crate) struct ProxySettingsSecret {
impl Store {
pub(crate) async fn installation_id(&self) -> Result<String> {
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?;
+7 -1
View File
@@ -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 {
+1
View File
@@ -54,6 +54,7 @@ impl Store {
}
pub async fn clear_statistics_storage(&self) -> Result<StatisticsStorage> {
let _write = self.writes.lock().await;
let mut transaction = self.pool.begin().await?;
sqlx::query("DELETE FROM llm_calls")
.execute(&mut *transaction)
+4 -1
View File
@@ -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<ToolCommit> {
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
+14
View File
@@ -0,0 +1,14 @@
use std::sync::Arc;
use tokio::sync::{Mutex, MutexGuard};
#[derive(Clone, Default)]
pub(crate) struct WriteCoordinator {
lock: Arc<Mutex<()>>,
}
impl WriteCoordinator {
pub(crate) async fn lock(&self) -> MutexGuard<'_, ()> {
self.lock.lock().await
}
}
+72 -1
View File
@@ -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()),
+42
View File
@@ -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(