refactor: cleanup pass on document embedding/retrieval
All checks were successful
CI / test (pull_request) Successful in 6m56s

- retrieval::retrieve_relevant now shows a spinner around the query-
  embedding call and gets its own tracing span, matching every other
  model-calling phase in the pipeline (it previously ran invisibly and
  untraced)
- add progress::report() to dedupe the show_progress-gated eprintln!
  pattern shared by research.rs and retrieval.rs
- documents::push_document builds its TextSplitter once per collect_documents
  call instead of once per file, skips the char-truncation walk entirely
  when a file is already under the limit, and pulls the source-string
  branch out of the map closure
- researcher::gather_findings drops a doc_context emptiness check that
  build_doc_context already guarantees
- documents.rs tests use tempfile::tempdir() instead of a hand-rolled
  TempDir type

Verified with cargo build/test/clippy -D warnings/fmt --check.
This commit is contained in:
Austin Schaefer 2026-08-19 15:13:07 +02:00
parent 53b4f572df
commit 19c92ccd52
8 changed files with 61 additions and 58 deletions

1
Cargo.lock generated
View file

@ -2129,6 +2129,7 @@ dependencies = [
"scraper", "scraper",
"serde", "serde",
"serde_json", "serde_json",
"tempfile",
"text-splitter", "text-splitter",
"tokio", "tokio",
"tracing", "tracing",

View file

@ -27,3 +27,4 @@ tracing-subscriber = { version = "0.3.23", features = ["env-filter"] }
[dev-dependencies] [dev-dependencies]
serde_json = "1" serde_json = "1"
tempfile = "3"

View file

@ -1,6 +1,6 @@
use rig::embeddings::{EmbedError, TextEmbedder}; use rig::embeddings::{EmbedError, TextEmbedder};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use text_splitter::TextSplitter; use text_splitter::{Characters, TextSplitter};
/// Safety valve against accidentally pointing `--doc` at a huge binary or /// Safety valve against accidentally pointing `--doc` at a huge binary or
/// log file — not a content limit. Real documents are chunked in full (see /// log file — not a content limit. Real documents are chunked in full (see
@ -40,6 +40,7 @@ impl rig::Embed for Document {
/// (permissions, non-UTF-8) are skipped with a warning rather than failing /// (permissions, non-UTF-8) are skipped with a warning rather than failing
/// the whole run. /// the whole run.
pub(crate) fn collect_documents(paths: &[PathBuf]) -> anyhow::Result<Vec<Document>> { pub(crate) fn collect_documents(paths: &[PathBuf]) -> anyhow::Result<Vec<Document>> {
let splitter = TextSplitter::new(CHUNK_CHARS);
let mut documents = Vec::new(); let mut documents = Vec::new();
for path in paths { for path in paths {
@ -47,30 +48,38 @@ pub(crate) fn collect_documents(paths: &[PathBuf]) -> anyhow::Result<Vec<Documen
for entry in std::fs::read_dir(path)? { for entry in std::fs::read_dir(path)? {
let entry = entry?; let entry = entry?;
if entry.file_type()?.is_file() { if entry.file_type()?.is_file() {
push_document(&entry.path(), &mut documents); push_document(&entry.path(), &splitter, &mut documents);
} }
} }
} else { } else {
push_document(path, &mut documents); push_document(path, &splitter, &mut documents);
} }
} }
Ok(documents) Ok(documents)
} }
fn push_document(path: &Path, documents: &mut Vec<Document>) { fn push_document(path: &Path, splitter: &TextSplitter<Characters>, documents: &mut Vec<Document>) {
match std::fs::read_to_string(path) { match std::fs::read_to_string(path) {
Ok(text) => { Ok(mut text) => {
let text: String = text.chars().take(MAX_FILE_CHARS).collect(); // Cheap upper bound on char count, so the vast majority of
let chunks: Vec<&str> = TextSplitter::new(CHUNK_CHARS).chunks(&text).collect(); // documents (well under the limit) skip the char-by-char walk
// entirely.
if text.len() > MAX_FILE_CHARS && text.chars().count() > MAX_FILE_CHARS {
text = text.chars().take(MAX_FILE_CHARS).collect();
}
let chunks: Vec<&str> = splitter.chunks(&text).collect();
let total = chunks.len(); let total = chunks.len();
let source = |i: usize| {
documents.extend(chunks.into_iter().enumerate().map(|(i, chunk)| Document { if total > 1 {
source: if total > 1 {
format!("{} (part {}/{total})", path.display(), i + 1) format!("{} (part {}/{total})", path.display(), i + 1)
} else { } else {
path.display().to_string() path.display().to_string()
}, }
};
documents.extend(chunks.into_iter().enumerate().map(|(i, chunk)| Document {
source: source(i),
text: chunk.to_string(), text: chunk.to_string(),
})); }));
} }
@ -83,40 +92,11 @@ fn push_document(path: &Path, documents: &mut Vec<Document>) {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use tempfile::tempdir;
/// A directory unique to this test run, under the OS temp dir, cleaned
/// up on drop so a panic mid-test can't leave stray files behind for the
/// next run.
struct TempDir(PathBuf);
impl TempDir {
fn new(name: &str) -> Self {
let path = std::env::temp_dir().join(format!(
"doubleo7-test-{name}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock is after the epoch")
.as_nanos()
));
std::fs::create_dir_all(&path).expect("create temp dir");
Self(path)
}
fn path(&self) -> &Path {
&self.0
}
}
impl Drop for TempDir {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
#[test] #[test]
fn collect_documents_reads_an_explicit_file() { fn collect_documents_reads_an_explicit_file() {
let dir = TempDir::new("explicit-file"); let dir = tempdir().unwrap();
let file = dir.path().join("notes.txt"); let file = dir.path().join("notes.txt");
std::fs::write(&file, "hello from a file").unwrap(); std::fs::write(&file, "hello from a file").unwrap();
@ -129,7 +109,7 @@ mod tests {
#[test] #[test]
fn collect_documents_reads_every_file_in_a_directory_non_recursively() { fn collect_documents_reads_every_file_in_a_directory_non_recursively() {
let dir = TempDir::new("directory"); let dir = tempdir().unwrap();
std::fs::write(dir.path().join("a.txt"), "a").unwrap(); std::fs::write(dir.path().join("a.txt"), "a").unwrap();
std::fs::write(dir.path().join("b.txt"), "b").unwrap(); std::fs::write(dir.path().join("b.txt"), "b").unwrap();
std::fs::create_dir(dir.path().join("nested")).unwrap(); std::fs::create_dir(dir.path().join("nested")).unwrap();
@ -150,7 +130,7 @@ mod tests {
#[test] #[test]
fn collect_documents_skips_unreadable_paths_instead_of_failing() { fn collect_documents_skips_unreadable_paths_instead_of_failing() {
let dir = TempDir::new("missing"); let dir = tempdir().unwrap();
let missing = dir.path().join("does-not-exist.txt"); let missing = dir.path().join("does-not-exist.txt");
assert!(collect_documents(&[missing]).unwrap().is_empty()); assert!(collect_documents(&[missing]).unwrap().is_empty());
@ -158,7 +138,7 @@ mod tests {
#[test] #[test]
fn collect_documents_splits_a_large_file_into_multiple_chunks() { fn collect_documents_splits_a_large_file_into_multiple_chunks() {
let dir = TempDir::new("large-file"); let dir = tempdir().unwrap();
let file = dir.path().join("big.txt"); let file = dir.path().join("big.txt");
// Well over CHUNK_CHARS, and with paragraph breaks so the splitter // Well over CHUNK_CHARS, and with paragraph breaks so the splitter
// has real semantic boundaries to chunk on. // has real semantic boundaries to chunk on.

View file

@ -270,7 +270,11 @@ mod tests {
serde_json::json!({ "query": "test" }), serde_json::json!({ "query": "test" }),
)], )],
}, },
Message::tool_result("call-1", "call-1-name", "1. Example\n https://example.com\n snippet"), Message::tool_result(
"call-1",
"call-1-name",
"1. Example\n https://example.com\n snippet",
),
Message::assistant("that source looks solid"), Message::assistant("that source looks solid"),
]; ];

View file

@ -69,6 +69,15 @@ impl Drop for Spinner {
} }
} }
/// Prints a status line, but only when progress display is on — matching
/// `Spinner`'s own `enabled` gate so a call site doesn't need to repeat the
/// `if show_progress` check inline.
pub(crate) fn report(enabled: bool, message: impl Into<String>) {
if enabled {
eprintln!("{}", message.into());
}
}
/// Updates the message of whatever spinner is currently running, if any. /// Updates the message of whatever spinner is currently running, if any.
/// A no-op when progress display is off (no spinner was ever started, so /// A no-op when progress display is off (no spinner was ever started, so
/// `active()` stays empty) or between phases (the previous `Spinner` has /// `active()` stays empty) or between phases (the previous `Spinner` has

View file

@ -1,5 +1,5 @@
use crate::documents; use crate::documents;
use crate::progress::REJECTED_EMOJI; use crate::progress::{self, REJECTED_EMOJI};
use crate::researcher::gather_findings; use crate::researcher::gather_findings;
use crate::retrieval; use crate::retrieval;
use crate::review::{self, Review}; use crate::review::{self, Review};
@ -57,12 +57,13 @@ pub(crate) async fn research(
break; break;
} }
if show_progress { progress::report(
eprintln!( show_progress,
format!(
"{REJECTED_EMOJI} Findings rejected — revising for round {}...", "{REJECTED_EMOJI} Findings rejected — revising for round {}...",
round + 1 round + 1
); ),
} );
feedback = Some(review); feedback = Some(review);
} }

View file

@ -70,14 +70,14 @@ pub(crate) async fn gather_findings(
}; };
let task = match doc_context { let task = match doc_context {
Some(context) if !context.is_empty() => format!( Some(context) => format!(
"{task}\n\n\ "{task}\n\n\
Relevant excerpts from documents the user uploaded treat these as trusted primary \ Relevant excerpts from documents the user uploaded treat these as trusted primary \
sources alongside anything you find on the web, and cite them with the same \ sources alongside anything you find on the web, and cite them with the same \
footnote scheme (their Sources entry can just be the document path shown below):\n\ footnote scheme (their Sources entry can just be the document path shown below):\n\
{context}" {context}"
), ),
_ => task, None => task,
}; };
let spinner = Spinner::start(show_progress, format!("{RESEARCH_EMOJI} Researching...")); let spinner = Spinner::start(show_progress, format!("{RESEARCH_EMOJI} Researching..."));

View file

@ -1,6 +1,6 @@
use crate::documents::Document; use crate::documents::Document;
use crate::models::EMBEDDING_MODEL; use crate::models::EMBEDDING_MODEL;
use crate::progress::{EMBED_EMOJI, Spinner}; use crate::progress::{self, EMBED_EMOJI, Spinner};
use rig::client::EmbeddingsClient; use rig::client::EmbeddingsClient;
use rig::providers::ollama; use rig::providers::ollama;
use rig::vector_store::VectorStoreIndex; use rig::vector_store::VectorStoreIndex;
@ -44,6 +44,7 @@ pub(crate) async fn build_index(
/// Retrieves the excerpts most semantically relevant to `query` from the /// Retrieves the excerpts most semantically relevant to `query` from the
/// index, formatted for inclusion in the researcher's task. /// index, formatted for inclusion in the researcher's task.
#[tracing::instrument(skip(index), fields(gen_ai.agent.name = "retriever"))]
pub(crate) async fn retrieve_relevant( pub(crate) async fn retrieve_relevant(
index: &DocumentIndex, index: &DocumentIndex,
query: &str, query: &str,
@ -54,14 +55,20 @@ pub(crate) async fn retrieve_relevant(
.samples(TOP_N_EXCERPTS) .samples(TOP_N_EXCERPTS)
.build(); .build();
let spinner = Spinner::start(
show_progress,
format!("{EMBED_EMOJI} Retrieving relevant excerpts..."),
);
let results = index.top_n::<Document>(request).await?; let results = index.top_n::<Document>(request).await?;
drop(spinner);
if show_progress { progress::report(
eprintln!( show_progress,
format!(
"{EMBED_EMOJI} Retrieved {} relevant excerpt(s) from uploaded documents", "{EMBED_EMOJI} Retrieved {} relevant excerpt(s) from uploaded documents",
results.len() results.len()
); ),
} );
tracing::info!(count = results.len(), "retrieved document excerpts"); tracing::info!(count = results.len(), "retrieved document excerpts");
Ok(results Ok(results