mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:40:50 +08:00
fix: harden concurrent persistence and task recovery
This commit is contained in:
@@ -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;
|
||||
});
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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
@@ -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
@@ -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>>> {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 = ?",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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()),
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user