refactor(config): consumer들이 &Config 대신 타입 슬라이스 수령 (god-struct 결합 해소)

This commit is contained in:
2026-06-24 11:55:10 +00:00
parent bf7769cf5a
commit 2dfbbe4f43
53 changed files with 310 additions and 257 deletions

View File

@@ -41,7 +41,7 @@ use kebab_core::{
Answer, DocumentStore, Embedder, ExtractContext, Extractor, IndexVersion, LanguageModel, Answer, DocumentStore, Embedder, ExtractContext, Extractor, IndexVersion, LanguageModel,
MediaType, Retriever, SearchHit, SearchMode, SearchOpts, SearchQuery, VectorStore, MediaType, Retriever, SearchHit, SearchMode, SearchOpts, SearchQuery, VectorStore,
}; };
use kebab_embed_local::FastembedEmbedder; use kebab_embed_local::{FASTEMBED_CACHE_SUBDIR, FastembedEmbedder};
use kebab_embed_ollama::OllamaEmbedder; use kebab_embed_ollama::OllamaEmbedder;
use kebab_llm_local::OllamaLanguageModel; use kebab_llm_local::OllamaLanguageModel;
use kebab_parse_code::{ use kebab_parse_code::{
@@ -138,7 +138,7 @@ impl App {
/// internally drives a `tokio::Runtime::block_on`, which panics if /// internally drives a `tokio::Runtime::block_on`, which panics if
/// invoked from inside another tokio runtime. /// invoked from inside another tokio runtime.
pub fn open_with_config(config: kebab_config::Config) -> Result<Self> { pub fn open_with_config(config: kebab_config::Config) -> Result<Self> {
let sqlite = SqliteStore::open(&config).context("kb-app: open SqliteStore")?; let sqlite = SqliteStore::open(&config.storage).context("kb-app: open SqliteStore")?;
sqlite sqlite
.run_migrations() .run_migrations()
.context("kb-app: run SqliteStore migrations")?; .context("kb-app: run SqliteStore migrations")?;
@@ -293,7 +293,7 @@ impl App {
vec_iv, vec_iv,
self.config.search.snippet_chars, self.config.search.snippet_chars,
)) as Arc<dyn Retriever>; )) as Arc<dyn Retriever>;
let hybrid = HybridRetriever::new(&self.config, lex, vec_retr); let hybrid = HybridRetriever::new(&self.config.search, lex, vec_retr);
hybrid.search(&query)? hybrid.search(&query)?
} }
}; };
@@ -391,7 +391,7 @@ impl App {
self.config.search.snippet_chars, self.config.search.snippet_chars,
)) as Arc<dyn Retriever> )) as Arc<dyn Retriever>
}; };
let hybrid = HybridRetriever::new(&self.config, lex, vec_retr); let hybrid = HybridRetriever::new(&self.config.search, lex, vec_retr);
let (mut traced_hits, trace) = hybrid.search_with_trace(&fetch_query)?; let (mut traced_hits, trace) = hybrid.search_with_trace(&fetch_query)?;
// Stamp staleness — same as search_uncached. // Stamp staleness — same as search_uncached.
@@ -535,7 +535,14 @@ impl App {
retriever: Arc<dyn Retriever>, retriever: Arc<dyn Retriever>,
llm: Arc<dyn LanguageModel>, llm: Arc<dyn LanguageModel>,
) -> RagPipeline { ) -> RagPipeline {
let pipeline = RagPipeline::new(self.config.clone(), retriever, llm, self.sqlite.clone()); let pipeline = RagPipeline::new(
self.config.rag.clone(),
self.config.models.clone(),
self.config.search.clone(),
retriever,
llm,
self.sqlite.clone(),
);
match &self.pipeline_verifier { match &self.pipeline_verifier {
Some(v) => pipeline.with_verifier(v.clone()), Some(v) => pipeline.with_verifier(v.clone()),
None => pipeline, None => pipeline,
@@ -583,7 +590,7 @@ impl App {
vec_iv, vec_iv,
self.config.search.snippet_chars, self.config.search.snippet_chars,
)) as Arc<dyn Retriever>; )) as Arc<dyn Retriever>;
Arc::new(HybridRetriever::new(&self.config, lex, vec_retr)) Arc::new(HybridRetriever::new(&self.config.search, lex, vec_retr))
} }
}) })
} }
@@ -613,12 +620,37 @@ impl App {
// offloads to a remote `/api/embed` daemon. // offloads to a remote `/api/embed` daemon.
let provider = self.config.models.embedding.provider.as_str(); let provider = self.config.models.embedding.provider.as_str();
let emb: Arc<dyn Embedder + Send + Sync> = match provider { let emb: Arc<dyn Embedder + Send + Sync> = match provider {
"fastembed" | "onnx" | "" => Arc::new( "fastembed" | "onnx" | "" => {
FastembedEmbedder::new(&self.config).context("kb-app: load FastembedEmbedder")?, // Resolve `{data_dir}/models/fastembed/` here so the
), // embedder constructor only takes the `[models.embedding]`
"ollama" => Arc::new( // slice + the final cache dir.
OllamaEmbedder::new(&self.config).context("kb-app: load OllamaEmbedder")?, let data_dir = kebab_config::expand_path(&self.config.storage.data_dir, "");
), let model_dir = kebab_config::expand_path(
&self.config.storage.model_dir,
&data_dir.to_string_lossy(),
);
let cache_dir = model_dir.join(FASTEMBED_CACHE_SUBDIR);
Arc::new(
FastembedEmbedder::new(&self.config.models.embedding, &cache_dir)
.context("kb-app: load FastembedEmbedder")?,
)
}
"ollama" => {
// Resolve the endpoint here: `models.embedding.endpoint`
// → fallback `models.llm.endpoint`.
let endpoint = self
.config
.models
.embedding
.endpoint
.clone()
.filter(|e| !e.is_empty())
.unwrap_or_else(|| self.config.models.llm.endpoint.clone());
Arc::new(
OllamaEmbedder::new(&self.config.models.embedding, endpoint)
.context("kb-app: load OllamaEmbedder")?,
)
}
other => { other => {
return Err(anyhow!( return Err(anyhow!(
"kb-app: unknown embedding provider {other:?}; expected one of \ "kb-app: unknown embedding provider {other:?}; expected one of \
@@ -643,7 +675,7 @@ impl App {
return Ok(Some(v.clone())); return Ok(Some(v.clone()));
} }
let store = Arc::new( let store = Arc::new(
LanceVectorStore::new(&self.config, self.sqlite.clone()) LanceVectorStore::new(&self.config.storage, self.sqlite.clone())
.context("kb-app: open LanceVectorStore")?, .context("kb-app: open LanceVectorStore")?,
); );
let _ = self.vector.set(store.clone()); let _ = self.vector.set(store.clone());
@@ -1054,7 +1086,7 @@ mod tests_trace {
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
// Bring up migrations. // Bring up migrations.
let store = kebab_store_sqlite::SqliteStore::open(&cfg).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
drop(store); drop(store);
let app = App::open_with_config(cfg).unwrap(); let app = App::open_with_config(cfg).unwrap();
@@ -1122,7 +1154,7 @@ mod tests_extractor_dispatch {
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
// Bring up migrations. // Bring up migrations.
let store = kebab_store_sqlite::SqliteStore::open(&cfg).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
drop(store); drop(store);
let app = App::open_with_config(cfg).unwrap(); let app = App::open_with_config(cfg).unwrap();

View File

@@ -262,7 +262,7 @@ mod tests {
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
// Bring up migrations so SqliteStore::open_existing succeeds inside App::open. // Bring up migrations so SqliteStore::open_existing succeeds inside App::open.
let store = kebab_store_sqlite::SqliteStore::open(&cfg).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
drop(store); drop(store);
// Leak the tempdir into a static — tests are short-lived; not worth threading. // Leak the tempdir into a static — tests are short-lived; not worth threading.

View File

@@ -139,7 +139,7 @@ pub fn enumerate_orphans(cfg: &Config) -> Result<Vec<WorkspacePath>> {
use kebab_core::SourceScope; use kebab_core::SourceScope;
use kebab_source_fs::FsSourceConnector; use kebab_source_fs::FsSourceConnector;
let store = kebab_store_sqlite::SqliteStore::open(cfg) let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage)
.context("enumerate_orphans: open SqliteStore")?; .context("enumerate_orphans: open SqliteStore")?;
let stored = store let stored = store
@@ -237,7 +237,7 @@ fn execute_orphans_only(cfg: &Config) -> Result<ResetReport> {
} }
let store = std::sync::Arc::new( let store = std::sync::Arc::new(
kebab_store_sqlite::SqliteStore::open(cfg) kebab_store_sqlite::SqliteStore::open(&cfg.storage)
.context("execute_orphans_only: open SqliteStore")?, .context("execute_orphans_only: open SqliteStore")?,
); );
@@ -296,7 +296,7 @@ fn open_vector_store_if_configured(
if cfg.models.embedding.provider == "none" || cfg.models.embedding.dimensions == 0 { if cfg.models.embedding.provider == "none" || cfg.models.embedding.dimensions == 0 {
return Ok(None); return Ok(None);
} }
match kebab_store_vector::LanceVectorStore::new(cfg, store) { match kebab_store_vector::LanceVectorStore::new(&cfg.storage, store) {
Ok(vs) => Ok(Some(vs)), Ok(vs) => Ok(Some(vs)),
Err(e) => { Err(e) => {
tracing::warn!( tracing::warn!(
@@ -320,7 +320,7 @@ fn truncate_embeddings(cfg: &Config) -> Result<u64> {
if !sqlite_path.exists() { if !sqlite_path.exists() {
return Ok(0); return Ok(0);
} }
let store = kebab_store_sqlite::SqliteStore::open(cfg) let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage)
.context("open SqliteStore for truncate_embedding_records")?; .context("open SqliteStore for truncate_embedding_records")?;
store.truncate_embedding_records() store.truncate_embedding_records()
} }

View File

@@ -263,7 +263,7 @@ mod tests_stats_ext {
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
// Bring up migrations so the sqlite file is created. // Bring up migrations so the sqlite file is created.
let store = kebab_store_sqlite::SqliteStore::open(&cfg).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
drop(store); drop(store);

View File

@@ -23,7 +23,7 @@ use kebab_core::{DocFilter, DocumentStore, SearchMode, SearchQuery, SourceScope}
/// Helper: open the store via `TestEnv` and run `list_documents`. /// Helper: open the store via `TestEnv` and run `list_documents`.
fn list_doc_paths(env: &TestEnv) -> Vec<String> { fn list_doc_paths(env: &TestEnv) -> Vec<String> {
use kebab_store_sqlite::SqliteStore; use kebab_store_sqlite::SqliteStore;
let store = SqliteStore::open(&env.config).unwrap(); let store = SqliteStore::open(&env.config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
.list_documents(&DocFilter::default()) .list_documents(&DocFilter::default())

View File

@@ -72,7 +72,7 @@ fn seed_ocr_events(env: &TestEnv, store: &SqliteStore) {
fn open_app_with_seeded_events(env: &TestEnv) -> App { fn open_app_with_seeded_events(env: &TestEnv) -> App {
let app = env.app(); let app = env.app();
let store = SqliteStore::open(&env.config).expect("open store for seed"); let store = SqliteStore::open(&env.config.storage).expect("open store for seed");
store.run_migrations().expect("run migrations for seed"); store.run_migrations().expect("run migrations for seed");
seed_ocr_events(env, &store); seed_ocr_events(env, &store);
app app

View File

@@ -23,7 +23,7 @@ use kebab_core::{DocFilter, DocumentStore, SourceScope};
/// Open the SqliteStore and list all `workspace_path` values. /// Open the SqliteStore and list all `workspace_path` values.
fn list_doc_paths(env: &TestEnv) -> Vec<String> { fn list_doc_paths(env: &TestEnv) -> Vec<String> {
use kebab_store_sqlite::SqliteStore; use kebab_store_sqlite::SqliteStore;
let store = SqliteStore::open(&env.config).unwrap(); let store = SqliteStore::open(&env.config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
.list_documents(&DocFilter::default()) .list_documents(&DocFilter::default())

View File

@@ -32,7 +32,7 @@ fn schema_models_active_arrays_empty_on_empty_corpus() {
std::fs::create_dir_all(&workspace).unwrap(); std::fs::create_dir_all(&workspace).unwrap();
let cfg = minimal_config(dir.path(), &workspace); let cfg = minimal_config(dir.path(), &workspace);
let store = kebab_store_sqlite::SqliteStore::open(&cfg).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
drop(store); drop(store);

View File

@@ -107,7 +107,7 @@ fn search_uncached_returns_same_hits_as_cached() {
#[test] #[test]
fn first_ingest_bumps_corpus_revision() { fn first_ingest_bumps_corpus_revision() {
let env = TestEnv::lexical_only(); let env = TestEnv::lexical_only();
let store_before = kebab_store_sqlite::SqliteStore::open(&env.config).unwrap(); let store_before = kebab_store_sqlite::SqliteStore::open(&env.config.storage).unwrap();
store_before.run_migrations().unwrap(); store_before.run_migrations().unwrap();
// V004 seeds 0; V009 + V010 + V011 migrations each bump by 1 to // V004 seeds 0; V009 + V010 + V011 migrations each bump by 1 to
// invalidate stale LRU caches (spec §5.2). Baseline before ingest = 3. // invalidate stale LRU caches (spec §5.2). Baseline before ingest = 3.
@@ -122,7 +122,7 @@ fn first_ingest_bumps_corpus_revision() {
"first ingest must commit ≥1 doc" "first ingest must commit ≥1 doc"
); );
let store_after = kebab_store_sqlite::SqliteStore::open(&env.config).unwrap(); let store_after = kebab_store_sqlite::SqliteStore::open(&env.config.storage).unwrap();
assert!( assert!(
store_after.corpus_revision() > baseline, store_after.corpus_revision() > baseline,
"ingest commit must bump corpus_revision past baseline {baseline} (got {})", "ingest commit must bump corpus_revision past baseline {baseline} (got {})",

View File

@@ -72,7 +72,7 @@ fn twin_files_fetch_span_uses_correct_asset() {
// Resolve doc_ids for both workspace paths. // Resolve doc_ids for both workspace paths.
// The ingest layer normalises workspace_path to the path relative to // The ingest layer normalises workspace_path to the path relative to
// workspace_root (e.g. "src_a/note.md"), so we look up by that form. // workspace_root (e.g. "src_a/note.md"), so we look up by that form.
let store = kebab_store_sqlite::SqliteStore::open(&env.config).unwrap(); let store = kebab_store_sqlite::SqliteStore::open(&env.config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// Find the twin items by matching on suffix so the test is robust to // Find the twin items by matching on suffix so the test is robust to

View File

@@ -23,17 +23,18 @@
//! See `docs/superpowers/specs/2026-04-27-kebab-final-form-design.md` //! See `docs/superpowers/specs/2026-04-27-kebab-final-form-design.md`
//! §7.2 (Embedder), §6.4 ([models.embedding]), §9 (versioning). //! §7.2 (Embedder), §6.4 ([models.embedding]), §9 (versioning).
use std::path::Path;
use std::sync::Mutex; use std::sync::Mutex;
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use fastembed::{EmbeddingModel, InitOptions, TextEmbedding}; use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
use kebab_config::expand_path; use kebab_config::EmbeddingModelCfg;
use kebab_embed::{Embedder, EmbeddingInput, EmbeddingKind, EmbeddingModelId, EmbeddingVersion}; use kebab_embed::{Embedder, EmbeddingInput, EmbeddingKind, EmbeddingModelId, EmbeddingVersion};
/// Subdirectory under `config.storage.model_dir` where the fastembed /// Subdirectory under `config.storage.model_dir` where the fastembed
/// adapter writes / reads ONNX + tokenizer files. Hard-coded per task /// adapter writes / reads ONNX + tokenizer files. Hard-coded per task
/// spec ("Model files cached under `config.storage.model_dir/fastembed/`"). /// spec ("Model files cached under `config.storage.model_dir/fastembed/`").
const FASTEMBED_CACHE_SUBDIR: &str = "fastembed"; pub const FASTEMBED_CACHE_SUBDIR: &str = "fastembed";
/// Local fastembed-rs adapter. /// Local fastembed-rs adapter.
/// ///
@@ -55,37 +56,35 @@ pub struct FastembedEmbedder {
} }
impl FastembedEmbedder { impl FastembedEmbedder {
/// Build an embedder from `Config`. Validates that /// Build an embedder from the `[models.embedding]` slice + a resolved
/// `config.models.embedding.dimensions` matches the model's actual /// `cache_dir` (the fastembed subdir under `config.storage.model_dir`;
/// dim BEFORE returning, so a mismatch fails at construction (not on /// the caller resolves it from the storage paths and the
/// first `embed`). /// [`FASTEMBED_CACHE_SUBDIR`] constant). Validates that `cfg.dimensions`
pub fn new(config: &kebab_config::Config) -> Result<Self> { /// matches the model's actual dim BEFORE returning, so a mismatch fails
// 1. Resolve `{data_dir}/models/fastembed/` from the config /// at construction (not on first `embed`).
// templates. Goes through the shared `kebab_config::expand_path` pub fn new(cfg: &EmbeddingModelCfg, cache_dir: &Path) -> Result<Self> {
// so every crate resolves storage paths identically. // 1. The caller resolved `{data_dir}/models/fastembed/`; we own
let data_dir = expand_path(&config.storage.data_dir, ""); // directory creation so a missing cache dir still works.
let model_dir = expand_path(&config.storage.model_dir, &data_dir.to_string_lossy()); std::fs::create_dir_all(cache_dir)
let cache_dir = model_dir.join(FASTEMBED_CACHE_SUBDIR);
std::fs::create_dir_all(&cache_dir)
.with_context(|| format!("create fastembed cache dir {}", cache_dir.display()))?; .with_context(|| format!("create fastembed cache dir {}", cache_dir.display()))?;
// 2. Resolve the fastembed enum variant from // 2. Resolve the fastembed enum variant from `cfg.model`. Currently
// `config.models.embedding.model`. Currently `multilingual-e5-large` // `multilingual-e5-large` (default) and `multilingual-e5-small`
// (default) and `multilingual-e5-small` are wired; other model names // are wired; other model names error out with a clear message
// error out with a clear message rather than silently misconfiguring. // rather than silently misconfiguring.
let model_name = resolve_model(&config.models.embedding.model)?; let model_name = resolve_model(&cfg.model)?;
// 3. Verify dim match BEFORE loading the model — if the config // 3. Verify dim match BEFORE loading the model — if the config
// is wrong we want to fail without paying the ONNX // is wrong we want to fail without paying the ONNX
// initialization cost. // initialization cost.
let model_info = let model_info =
TextEmbedding::get_model_info(&model_name).context("fastembed: get_model_info")?; TextEmbedding::get_model_info(&model_name).context("fastembed: get_model_info")?;
check_dim(model_info.dim, config.models.embedding.dimensions)?; check_dim(model_info.dim, cfg.dimensions)?;
tracing::info!( tracing::info!(
target: "kebab-embed-local", target: "kebab-embed-local",
cache_dir = %cache_dir.display(), cache_dir = %cache_dir.display(),
model = %config.models.embedding.model, model = %cfg.model,
dims = model_info.dim, dims = model_info.dim,
"initializing FastembedEmbedder" "initializing FastembedEmbedder"
); );
@@ -95,11 +94,11 @@ impl FastembedEmbedder {
// download progress is surfaced via the `tracing::info!` // download progress is surfaced via the `tracing::info!`
// pair around `TextEmbedding::try_new` instead. // pair around `TextEmbedding::try_new` instead.
let opts = InitOptions::new(model_name.clone()) let opts = InitOptions::new(model_name.clone())
.with_cache_dir(cache_dir.clone()) .with_cache_dir(cache_dir.to_path_buf())
.with_show_download_progress(false); .with_show_download_progress(false);
tracing::info!( tracing::info!(
target: "kebab-embed-local", target: "kebab-embed-local",
model = %config.models.embedding.model, model = %cfg.model,
cache_dir = %cache_dir.display(), cache_dir = %cache_dir.display(),
"loading embedding model (first run downloads model weights — ~470MB for e5-small, ~1.3GB for e5-large)" "loading embedding model (first run downloads model weights — ~470MB for e5-small, ~1.3GB for e5-large)"
); );
@@ -107,17 +106,17 @@ impl FastembedEmbedder {
let dimensions = model_info.dim; let dimensions = model_info.dim;
tracing::info!( tracing::info!(
target: "kebab-embed-local", target: "kebab-embed-local",
model = %config.models.embedding.model, model = %cfg.model,
dimensions, dimensions,
"embedding model loaded" "embedding model loaded"
); );
Ok(Self { Ok(Self {
inner: Mutex::new(inner), inner: Mutex::new(inner),
model_id: EmbeddingModelId(config.models.embedding.model.clone()), model_id: EmbeddingModelId(cfg.model.clone()),
version: EmbeddingVersion(config.models.embedding.version.clone()), version: EmbeddingVersion(cfg.version.clone()),
dimensions, dimensions,
batch_size: config.models.embedding.batch_size, batch_size: cfg.batch_size,
}) })
} }
} }

View File

@@ -24,7 +24,15 @@ use std::sync::OnceLock;
use std::time::Instant; use std::time::Instant;
use kebab_embed::{Embedder, EmbeddingInput, EmbeddingKind}; use kebab_embed::{Embedder, EmbeddingInput, EmbeddingKind};
use kebab_embed_local::FastembedEmbedder; use kebab_embed_local::{FASTEMBED_CACHE_SUBDIR, FastembedEmbedder};
/// Resolve the fastembed cache dir from a `Config`'s storage paths,
/// mirroring what `kebab-app`'s `embedder()` does at the call site.
fn fastembed_cache_dir(cfg: &kebab_config::Config) -> std::path::PathBuf {
let data_dir = kebab_config::expand_path(&cfg.storage.data_dir, "");
let model_dir = kebab_config::expand_path(&cfg.storage.model_dir, &data_dir.to_string_lossy());
model_dir.join(FASTEMBED_CACHE_SUBDIR)
}
/// Build a `Config` whose `data_dir` lives in a per-process temp dir so /// Build a `Config` whose `data_dir` lives in a per-process temp dir so
/// the test never writes into the developer's real `~/.local/share/kebab`. /// the test never writes into the developer's real `~/.local/share/kebab`.
@@ -52,7 +60,8 @@ fn shared_embedder() -> &'static FastembedEmbedder {
// and wreck subsequent calls.) The OS will reclaim the leaked // and wreck subsequent calls.) The OS will reclaim the leaked
// path when the test process exits. // path when the test process exits.
let _ = std::mem::ManuallyDrop::new(_tmp); let _ = std::mem::ManuallyDrop::new(_tmp);
FastembedEmbedder::new(&cfg).expect("init FastembedEmbedder") let cache_dir = fastembed_cache_dir(&cfg);
FastembedEmbedder::new(&cfg.models.embedding, &cache_dir).expect("init FastembedEmbedder")
}) })
} }
@@ -73,10 +82,11 @@ fn default_config_constructs_with_dims_1024() {
fn mismatched_dims_in_config_errors_at_construction() { fn mismatched_dims_in_config_errors_at_construction() {
let (mut cfg, _tmp) = test_config(); let (mut cfg, _tmp) = test_config();
cfg.models.embedding.dimensions = 512; // model is 1024 (e5-large default) cfg.models.embedding.dimensions = 512; // model is 1024 (e5-large default)
let cache_dir = fastembed_cache_dir(&cfg);
// `FastembedEmbedder` deliberately does not implement `Debug` // `FastembedEmbedder` deliberately does not implement `Debug`
// (its inner ONNX session has no useful debug shape), so we // (its inner ONNX session has no useful debug shape), so we
// can't use `expect_err`; match the Result manually. // can't use `expect_err`; match the Result manually.
let err = match FastembedEmbedder::new(&cfg) { let err = match FastembedEmbedder::new(&cfg.models.embedding, &cache_dir) {
Ok(_) => panic!("dim mismatch must error"), Ok(_) => panic!("dim mismatch must error"),
Err(e) => e, Err(e) => e,
}; };

View File

@@ -43,6 +43,7 @@
use std::time::Duration; use std::time::Duration;
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use kebab_config::EmbeddingModelCfg;
use kebab_core::{Embedder, EmbeddingInput, EmbeddingKind, EmbeddingModelId, EmbeddingVersion}; use kebab_core::{Embedder, EmbeddingInput, EmbeddingKind, EmbeddingModelId, EmbeddingVersion};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
@@ -101,19 +102,15 @@ pub struct OllamaEmbedder {
} }
impl OllamaEmbedder { impl OllamaEmbedder {
/// Build from a workspace [`kebab_config::Config`]. Reads /// Build from the `[models.embedding]` slice + a resolved `endpoint`.
/// `config.models.embedding.{model, dimensions}` and resolves the endpoint /// Reads `cfg.{model, dimensions}`; the caller resolves the endpoint
/// as `models.embedding.endpoint` → fallback `models.llm.endpoint`. /// (`models.embedding.endpoint` → fallback `models.llm.endpoint`) and
/// passes it in.
/// ///
/// Does NOT touch the network. The caller (app layer) is expected to have /// Does NOT touch the network. The caller (app layer) is expected to have
/// validated `provider == "ollama"`. /// validated `provider == "ollama"`.
pub fn new(config: &kebab_config::Config) -> Result<Self> { pub fn new(cfg: &EmbeddingModelCfg, endpoint: String) -> Result<Self> {
let emb = &config.models.embedding; let emb = cfg;
let endpoint = emb
.endpoint
.clone()
.filter(|e| !e.is_empty())
.unwrap_or_else(|| config.models.llm.endpoint.clone());
if endpoint.is_empty() { if endpoint.is_empty() {
anyhow::bail!( anyhow::bail!(
"ollama embedding provider needs an endpoint: set \ "ollama embedding provider needs an endpoint: set \

View File

@@ -27,7 +27,16 @@ async fn embed_blocking(
inputs: Vec<(String, EmbeddingKind)>, inputs: Vec<(String, EmbeddingKind)>,
) -> anyhow::Result<Vec<Vec<f32>>> { ) -> anyhow::Result<Vec<Vec<f32>>> {
tokio::task::spawn_blocking(move || -> anyhow::Result<Vec<Vec<f32>>> { tokio::task::spawn_blocking(move || -> anyhow::Result<Vec<Vec<f32>>> {
let emb = OllamaEmbedder::new(&cfg)?; // Resolve the endpoint exactly as kebab-app's `embedder()` does:
// `models.embedding.endpoint` → fallback `models.llm.endpoint`.
let endpoint = cfg
.models
.embedding
.endpoint
.clone()
.filter(|e| !e.is_empty())
.unwrap_or_else(|| cfg.models.llm.endpoint.clone());
let emb = OllamaEmbedder::new(&cfg.models.embedding, endpoint)?;
let refs: Vec<EmbeddingInput<'_>> = inputs let refs: Vec<EmbeddingInput<'_>> = inputs
.iter() .iter()
.map(|(t, k)| EmbeddingInput { text: t, kind: *k }) .map(|(t, k)| EmbeddingInput { text: t, kind: *k })

View File

@@ -91,7 +91,7 @@ pub fn compare_runs_with_config(
run_id_b: &str, run_id_b: &str,
opts: &CompareOpts, opts: &CompareOpts,
) -> Result<CompareReport> { ) -> Result<CompareReport> {
let store = SqliteStore::open(cfg).context("open SqliteStore for compare_runs")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for compare_runs")?;
store.run_migrations().context("run migrations")?; store.run_migrations().context("run migrations")?;
// Pull both run rows up-front so we can extract chunker_version and // Pull both run rows up-front so we can extract chunker_version and

View File

@@ -127,7 +127,7 @@ pub(crate) fn validate_against_db(
return Ok(()); return Ok(());
} }
let store = SqliteStore::open(cfg).context("open SqliteStore for golden validation")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for golden validation")?;
store store
.run_migrations() .run_migrations()
.context("run migrations for golden validation")?; .context("run migrations for golden validation")?;
@@ -232,7 +232,7 @@ mod tests {
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = tmp.path().to_string_lossy().into_owned(); config.storage.data_dir = tmp.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&config).unwrap(); let store = SqliteStore::open(&config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
seed_one_chunk(&store, "doc_present", "chunk_present"); seed_one_chunk(&store, "doc_present", "chunk_present");
@@ -256,7 +256,7 @@ mod tests {
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = tmp.path().to_string_lossy().into_owned(); config.storage.data_dir = tmp.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&config).unwrap(); let store = SqliteStore::open(&config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
seed_one_chunk(&store, "doc_present", "chunk_present"); seed_one_chunk(&store, "doc_present", "chunk_present");

View File

@@ -114,7 +114,7 @@ pub fn compute_aggregate(run_id: &str) -> Result<AggregateMetrics> {
/// Compute aggregate metrics for `run_id` against an explicit /// Compute aggregate metrics for `run_id` against an explicit
/// [`Config`] (used by tests with a TempDir-backed `data_dir`). /// [`Config`] (used by tests with a TempDir-backed `data_dir`).
pub fn compute_aggregate_with_config(cfg: &Config, run_id: &str) -> Result<AggregateMetrics> { pub fn compute_aggregate_with_config(cfg: &Config, run_id: &str) -> Result<AggregateMetrics> {
let store = SqliteStore::open(cfg).context("open SqliteStore for compute_aggregate")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for compute_aggregate")?;
store store
.run_migrations() .run_migrations()
.context("run migrations for compute_aggregate")?; .context("run migrations for compute_aggregate")?;
@@ -146,7 +146,7 @@ pub fn store_aggregate_with_config(
run_id: &str, run_id: &str,
agg: &AggregateMetrics, agg: &AggregateMetrics,
) -> Result<()> { ) -> Result<()> {
let store = SqliteStore::open(cfg).context("open SqliteStore for store_aggregate")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for store_aggregate")?;
store.run_migrations().context("run migrations")?; store.run_migrations().context("run migrations")?;
let json = serde_json::to_string(agg).context("serialize AggregateMetrics")?; let json = serde_json::to_string(agg).context("serialize AggregateMetrics")?;
store store

View File

@@ -61,7 +61,7 @@ pub fn run_eval_with_config(cfg: &kebab_config::Config, opts: &EvalRunOpts) -> R
// Open the store once so every per-query write reuses the same // Open the store once so every per-query write reuses the same
// connection-mutex lifetime. // connection-mutex lifetime.
let store = SqliteStore::open(cfg).context("open SqliteStore for run_eval")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for run_eval")?;
store store
.run_migrations() .run_migrations()
.context("run migrations for run_eval")?; .context("run migrations for run_eval")?;

View File

@@ -239,7 +239,7 @@ pub fn compute_variant_consistency_with_config(
cfg: &Config, cfg: &Config,
run_id: &str, run_id: &str,
) -> Result<VariantConsistencyReport> { ) -> Result<VariantConsistencyReport> {
let store = SqliteStore::open(cfg).context("open SqliteStore for variant consistency")?; let store = SqliteStore::open(&cfg.storage).context("open SqliteStore for variant consistency")?;
store.run_migrations().context("run migrations")?; store.run_migrations().context("run migrations")?;
let run_record = store let run_record = store
.load_eval_run(run_id) .load_eval_run(run_id)

View File

@@ -152,7 +152,7 @@ fn compute_and_store_aggregate_round_trips() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
write_run( write_run(
@@ -183,7 +183,7 @@ fn compute_and_store_aggregate_round_trips() {
assert_eq!(agg.mrr, 0.4167); assert_eq!(agg.mrr, 0.4167);
store_aggregate_with_config(&cfg, "run_a", &agg).unwrap(); store_aggregate_with_config(&cfg, "run_a", &agg).unwrap();
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
let row = store.load_eval_run("run_a").unwrap().unwrap(); let row = store.load_eval_run("run_a").unwrap().unwrap();
let parsed: AggregateMetrics = serde_json::from_str(&row.aggregate_json).unwrap(); let parsed: AggregateMetrics = serde_json::from_str(&row.aggregate_json).unwrap();
// f32 round-trip via JSON is exact for our 4-decimal-rounded // f32 round-trip via JSON is exact for our 4-decimal-rounded
@@ -224,7 +224,7 @@ fn compare_runs_classifies_win_loss_draw_regression() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
// Run A: // Run A:
@@ -284,7 +284,7 @@ fn compare_strict_mode_refuses_chunker_version_mismatch() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
write_run( write_run(
@@ -316,7 +316,7 @@ fn compare_graceful_falls_back_to_doc_id() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
// Run A uses test@1 chunker; run B uses test@2 — chunk_ids no longer // Run A uses test@1 chunker; run B uses test@2 — chunk_ids no longer
@@ -357,7 +357,7 @@ fn compare_report_snapshot_matches_fixture() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
write_run( write_run(
@@ -434,7 +434,7 @@ fn render_report_md_is_human_readable() {
let _g = env_guard(); let _g = env_guard();
let tmp = TempDir::new().unwrap(); let tmp = TempDir::new().unwrap();
let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic()); let cfg = cfg_with_data_dir(&tmp, golden_yaml_basic());
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let now = OffsetDateTime::UNIX_EPOCH; let now = OffsetDateTime::UNIX_EPOCH;
write_run( write_run(

View File

@@ -48,7 +48,7 @@ impl RunEnv {
// Pin search defaults so test asserts are stable. // Pin search defaults so test asserts are stable.
config.search.default_k = 5; config.search.default_k = 5;
let store = SqliteStore::open(&config).unwrap(); let store = SqliteStore::open(&config.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
seed_corpus(&store); seed_corpus(&store);
Self { temp, config } Self { temp, config }
@@ -273,7 +273,7 @@ fn runner_persists_eval_run_and_query_result_rows() {
// the rows back. We use the inherent `read_conn` helper rather // the rows back. We use the inherent `read_conn` helper rather
// than rusqlite directly because the latter would require kb-eval // than rusqlite directly because the latter would require kb-eval
// to add a runtime rusqlite dep (forbidden by the spec). // to add a runtime rusqlite dep (forbidden by the spec).
let store = SqliteStore::open(&env.config).unwrap(); let store = SqliteStore::open(&env.config.storage).unwrap();
let conn = store.read_conn(); let conn = store.read_conn();
let n_runs: i64 = conn let n_runs: i64 = conn

View File

@@ -173,7 +173,14 @@ impl Default for AskOpts {
/// Single-threaded RAG orchestrator. See module docs for the stage list. /// Single-threaded RAG orchestrator. See module docs for the stage list.
pub struct RagPipeline { pub struct RagPipeline {
config: kebab_config::Config, /// `[rag]` policy slice (score gate, prompt template, multi-hop knobs,
/// NLI threshold). Replaces the old whole-`Config` field.
rag: kebab_config::RagCfg,
/// `[models]` slice — only `llm.temperature` / `llm.seed` and the
/// `embedding` block (via [`embedding_ref_for`]) are read.
models: kebab_config::ModelsCfg,
/// `[search]` slice — only `default_k` + `stale_threshold_days` read.
search: kebab_config::SearchCfg,
retriever: Arc<dyn Retriever>, retriever: Arc<dyn Retriever>,
llm: Arc<dyn LanguageModel>, llm: Arc<dyn LanguageModel>,
docs: Arc<SqliteStore>, docs: Arc<SqliteStore>,
@@ -192,16 +199,20 @@ impl RagPipeline {
/// inject mocks). /// inject mocks).
/// ///
/// The NLI verifier is NOT a constructor arg — it threads in via /// The NLI verifier is NOT a constructor arg — it threads in via
/// the [`Self::with_verifier`] builder so the historical 4-arg /// the [`Self::with_verifier`] builder so the verifier stays
/// signature stays stable across the PR-9c-1 surface bump. /// orthogonal to the core slice args.
pub fn new( pub fn new(
config: kebab_config::Config, rag: kebab_config::RagCfg,
models: kebab_config::ModelsCfg,
search: kebab_config::SearchCfg,
retriever: Arc<dyn Retriever>, retriever: Arc<dyn Retriever>,
llm: Arc<dyn LanguageModel>, llm: Arc<dyn LanguageModel>,
docs: Arc<SqliteStore>, docs: Arc<SqliteStore>,
) -> Self { ) -> Self {
Self { Self {
config, rag,
models,
search,
retriever, retriever,
llm, llm,
docs, docs,
@@ -237,7 +248,7 @@ impl RagPipeline {
// ── 1. Retrieve ──────────────────────────────────────────────────── // ── 1. Retrieve ────────────────────────────────────────────────────
// floor at config default — see `AskOpts::k` doc for rationale. // floor at config default — see `AskOpts::k` doc for rationale.
let k_effective = opts.k.max(self.config.search.default_k); let k_effective = opts.k.max(self.search.default_k);
let search_query = SearchQuery { let search_query = SearchQuery {
text: query.to_string(), text: query.to_string(),
mode: opts.mode, mode: opts.mode,
@@ -254,7 +265,7 @@ impl RagPipeline {
// `hit.stale` downstream, so stamping once here keeps both // `hit.stale` downstream, so stamping once here keeps both
// call sites aligned with the App-level `search` post-process. // call sites aligned with the App-level `search` post-process.
let now = OffsetDateTime::now_utc(); let now = OffsetDateTime::now_utc();
let stale_threshold_days = self.config.search.stale_threshold_days; let stale_threshold_days = self.search.stale_threshold_days;
for h in &mut hits { for h in &mut hits {
h.stale = compute_stale(h.indexed_at, now, stale_threshold_days); h.stale = compute_stale(h.indexed_at, now, stale_threshold_days);
} }
@@ -282,7 +293,7 @@ impl RagPipeline {
if hits.is_empty() { if hits.is_empty() {
return self.refuse_no_chunks(query, &opts, k_effective, started, None); return self.refuse_no_chunks(query, &opts, k_effective, started, None);
} }
if top_score < self.config.rag.score_gate { if top_score < self.rag.score_gate {
return self.refuse_score_gate(query, &opts, &hits, k_effective, started, None); return self.refuse_score_gate(query, &opts, &hits, k_effective, started, None);
} }
@@ -305,7 +316,7 @@ impl RagPipeline {
} }
// ── 4. Render prompt ─────────────────────────────────────────────── // ── 4. Render prompt ───────────────────────────────────────────────
let system = system_prompt_for(&self.config.rag.prompt_template_version)?.to_string(); let system = system_prompt_for(&self.rag.prompt_template_version)?.to_string();
let user = format!("[질문]\n{query}\n\n[근거]\n{packed_text}"); let user = format!("[질문]\n{query}\n\n[근거]\n{packed_text}");
// ── 5. Generate ──────────────────────────────────────────────────── // ── 5. Generate ────────────────────────────────────────────────────
@@ -321,8 +332,8 @@ impl RagPipeline {
let max_completion = llm_ctx.saturating_sub(used_for_input).max(64); let max_completion = llm_ctx.saturating_sub(used_for_input).max(64);
let temperature = opts let temperature = opts
.temperature .temperature
.unwrap_or(self.config.models.llm.temperature); .unwrap_or(self.models.llm.temperature);
let seed = opts.seed.or(Some(self.config.models.llm.seed)); let seed = opts.seed.or(Some(self.models.llm.seed));
let req = GenerateRequest { let req = GenerateRequest {
system: system.clone(), system: system.clone(),
user: user.clone(), user: user.clone(),
@@ -440,7 +451,7 @@ impl RagPipeline {
}) })
.collect(); .collect();
let embedding_ref = embedding_ref_for(opts.mode, &self.config); let embedding_ref = embedding_ref_for(opts.mode, &self.models);
let trace_id = mint_trace_id(query, top_score, &self.llm.model_ref().id); let trace_id = mint_trace_id(query, top_score, &self.llm.model_ref().id);
@@ -466,13 +477,13 @@ impl RagPipeline {
model: self.llm.model_ref(), model: self.llm.model_ref(),
embedding: embedding_ref, embedding: embedding_ref,
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
self.config.rag.prompt_template_version.clone(), self.rag.prompt_template_version.clone(),
), ),
retrieval: AnswerRetrievalSummary { retrieval: AnswerRetrievalSummary {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: k_effective, k: k_effective,
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score, top_score,
chunks_returned, chunks_returned,
chunks_used, chunks_used,
@@ -570,7 +581,7 @@ impl RagPipeline {
/// eval `compare` can isolate multi-hop runs from single-pass. /// eval `compare` can isolate multi-hop runs from single-pass.
pub fn ask_multi_hop(&self, query: &str, opts: AskOpts) -> Result<Answer> { pub fn ask_multi_hop(&self, query: &str, opts: AskOpts) -> Result<Answer> {
let started = std::time::Instant::now(); let started = std::time::Instant::now();
let k_effective = opts.k.max(self.config.search.default_k); let k_effective = opts.k.max(self.search.default_k);
// ── 0. Pre-decompose score-gate probe (v0.18 dogfood fix) ────────── // ── 0. Pre-decompose score-gate probe (v0.18 dogfood fix) ──────────
// //
@@ -606,14 +617,14 @@ impl RagPipeline {
.search(&probe_query) .search(&probe_query)
.context("kb-rag: multi-hop probe retriever.search")?; .context("kb-rag: multi-hop probe retriever.search")?;
let probe_now = OffsetDateTime::now_utc(); let probe_now = OffsetDateTime::now_utc();
let probe_threshold = self.config.search.stale_threshold_days; let probe_threshold = self.search.stale_threshold_days;
for h in &mut probe_hits { for h in &mut probe_hits {
h.stale = compute_stale(h.indexed_at, probe_now, probe_threshold); h.stale = compute_stale(h.indexed_at, probe_now, probe_threshold);
} }
if probe_hits.is_empty() { if probe_hits.is_empty() {
return self.refuse_no_chunks(query, &opts, k_effective, started, None); return self.refuse_no_chunks(query, &opts, k_effective, started, None);
} }
if probe_hits[0].retrieval.fusion_score < self.config.rag.score_gate { if probe_hits[0].retrieval.fusion_score < self.rag.score_gate {
return self.refuse_score_gate(query, &opts, &probe_hits, k_effective, started, None); return self.refuse_score_gate(query, &opts, &probe_hits, k_effective, started, None);
} }
@@ -658,8 +669,8 @@ impl RagPipeline {
// (stop); the loop also breaks when `max_depth` or // (stop); the loop also breaks when `max_depth` or
// `max_pool_chunks` cap fires (`forced_stop = true`). // `max_pool_chunks` cap fires (`forced_stop = true`).
// `k_effective` already computed at the probe step above. // `k_effective` already computed at the probe step above.
let max_depth = self.config.rag.multi_hop_max_depth; let max_depth = self.rag.multi_hop_max_depth;
let max_pool = self.config.rag.multi_hop_max_pool_chunks as usize; let max_pool = self.rag.multi_hop_max_pool_chunks as usize;
let mut pool: Vec<SearchHit> = Vec::new(); let mut pool: Vec<SearchHit> = Vec::new();
let mut seen_chunk_ids: std::collections::HashSet<String> = let mut seen_chunk_ids: std::collections::HashSet<String> =
std::collections::HashSet::new(); std::collections::HashSet::new();
@@ -754,7 +765,7 @@ impl RagPipeline {
// single-pass `hits` from here on — score gate / no-chunks / // single-pass `hits` from here on — score gate / no-chunks /
// pack_context all read it the same way. // pack_context all read it the same way.
let now = OffsetDateTime::now_utc(); let now = OffsetDateTime::now_utc();
let stale_threshold_days = self.config.search.stale_threshold_days; let stale_threshold_days = self.search.stale_threshold_days;
for h in &mut pool { for h in &mut pool {
h.stale = compute_stale(h.indexed_at, now, stale_threshold_days); h.stale = compute_stale(h.indexed_at, now, stale_threshold_days);
} }
@@ -775,7 +786,7 @@ impl RagPipeline {
if pool.is_empty() { if pool.is_empty() {
return self.refuse_no_chunks(query, &opts, k_effective, started, Some(hops)); return self.refuse_no_chunks(query, &opts, k_effective, started, Some(hops));
} }
if top_score < self.config.rag.score_gate { if top_score < self.rag.score_gate {
return self.refuse_score_gate(query, &opts, &pool, k_effective, started, Some(hops)); return self.refuse_score_gate(query, &opts, &pool, k_effective, started, Some(hops));
} }
@@ -816,8 +827,8 @@ impl RagPipeline {
let max_completion = llm_ctx.saturating_sub(used_for_input).max(64); let max_completion = llm_ctx.saturating_sub(used_for_input).max(64);
let temperature = opts let temperature = opts
.temperature .temperature
.unwrap_or(self.config.models.llm.temperature); .unwrap_or(self.models.llm.temperature);
let seed = opts.seed.or(Some(self.config.models.llm.seed)); let seed = opts.seed.or(Some(self.models.llm.seed));
let req = GenerateRequest { let req = GenerateRequest {
system: system.clone(), system: system.clone(),
user: user.clone(), user: user.clone(),
@@ -909,7 +920,7 @@ impl RagPipeline {
// (LlmStreamAborted) above; skipping the NLI gate here avoids // (LlmStreamAborted) above; skipping the NLI gate here avoids
// tokenizing an empty hypothesis (degenerate CLS-SEP-SEP that // tokenizing an empty hypothesis (degenerate CLS-SEP-SEP that
// would yield a near-uniform softmax and a misleading nli_passed). // would yield a near-uniform softmax and a misleading nli_passed).
let verification = if self.config.rag.nli_threshold > 0.0 && !acc.trim().is_empty() { let verification = if self.rag.nli_threshold > 0.0 && !acc.trim().is_empty() {
let v = self.verifier.as_ref().expect( let v = self.verifier.as_ref().expect(
"verifier must be Some when nli_threshold > 0.0 \ "verifier must be Some when nli_threshold > 0.0 \
(kebab-app's open_with_config enforces this invariant)", (kebab-app's open_with_config enforces this invariant)",
@@ -946,10 +957,10 @@ impl RagPipeline {
} }
match v.score(&truncated_premise, &truncated_hypothesis) { match v.score(&truncated_premise, &truncated_hypothesis) {
Ok(scores) => { Ok(scores) => {
let passed = scores.entailment >= self.config.rag.nli_threshold; let passed = scores.entailment >= self.rag.nli_threshold;
Some(VerificationSummary { Some(VerificationSummary {
nli_score: scores.entailment, nli_score: scores.entailment,
nli_threshold: self.config.rag.nli_threshold, nli_threshold: self.rag.nli_threshold,
nli_passed: passed, nli_passed: passed,
}) })
} }
@@ -984,7 +995,7 @@ impl RagPipeline {
}) })
.collect(); .collect();
let embedding_ref = embedding_ref_for(opts.mode, &self.config); let embedding_ref = embedding_ref_for(opts.mode, &self.models);
let trace_id = mint_trace_id(query, top_score, &self.llm.model_ref().id); let trace_id = mint_trace_id(query, top_score, &self.llm.model_ref().id);
let chunks_used = u32::try_from(packed_entries.len()).unwrap_or(u32::MAX); let chunks_used = u32::try_from(packed_entries.len()).unwrap_or(u32::MAX);
let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX); let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX);
@@ -1025,7 +1036,7 @@ impl RagPipeline {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: k_effective, k: k_effective,
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score, top_score,
chunks_returned, chunks_returned,
chunks_used, chunks_used,
@@ -1102,7 +1113,7 @@ impl RagPipeline {
query: &str, query: &str,
opts: &AskOpts, opts: &AskOpts,
) -> Result<(Option<Vec<String>>, u32)> { ) -> Result<(Option<Vec<String>>, u32)> {
let max = self.config.rag.multi_hop_max_sub_queries_per_iter as usize; let max = self.rag.multi_hop_max_sub_queries_per_iter as usize;
// `format!` named args give compile-time substitution checking // `format!` named args give compile-time substitution checking
// (PR-2 회차 1 carry-over fix): a typo in the template aborts // (PR-2 회차 1 carry-over fix): a typo in the template aborts
// compilation rather than silently emitting an unsubstituted // compilation rather than silently emitting an unsubstituted
@@ -1112,8 +1123,8 @@ impl RagPipeline {
); );
let temperature = opts let temperature = opts
.temperature .temperature
.unwrap_or(self.config.models.llm.temperature); .unwrap_or(self.models.llm.temperature);
let seed = opts.seed.or(Some(self.config.models.llm.seed)); let seed = opts.seed.or(Some(self.models.llm.seed));
let req = GenerateRequest { let req = GenerateRequest {
system: MULTI_HOP_DECOMPOSE_SYSTEM_PROMPT.to_string(), system: MULTI_HOP_DECOMPOSE_SYSTEM_PROMPT.to_string(),
user, user,
@@ -1172,14 +1183,14 @@ impl RagPipeline {
depth_remaining: u32, depth_remaining: u32,
opts: &AskOpts, opts: &AskOpts,
) -> Result<(Option<Vec<String>>, u32)> { ) -> Result<(Option<Vec<String>>, u32)> {
let max = self.config.rag.multi_hop_max_sub_queries_per_iter as usize; let max = self.rag.multi_hop_max_sub_queries_per_iter as usize;
let user = format!( let user = format!(
"[원본 질문]\n{query}\n\n[지금까지 모은 근거] ({pool_size} chunks)\n{packed_context}\n\n남은 깊이: {depth_remaining}\n\n추가 retrieval 이 필요하면 새 sub-question 들 (최대 {max} 개) 을 JSON array of strings 로, 충분하면 빈 array `[]` 를 반환:", "[원본 질문]\n{query}\n\n[지금까지 모은 근거] ({pool_size} chunks)\n{packed_context}\n\n남은 깊이: {depth_remaining}\n\n추가 retrieval 이 필요하면 새 sub-question 들 (최대 {max} 개) 을 JSON array of strings 로, 충분하면 빈 array `[]` 를 반환:",
); );
let temperature = opts let temperature = opts
.temperature .temperature
.unwrap_or(self.config.models.llm.temperature); .unwrap_or(self.models.llm.temperature);
let seed = opts.seed.or(Some(self.config.models.llm.seed)); let seed = opts.seed.or(Some(self.models.llm.seed));
let req = GenerateRequest { let req = GenerateRequest {
system: MULTI_HOP_DECIDE_SYSTEM_PROMPT.to_string(), system: MULTI_HOP_DECIDE_SYSTEM_PROMPT.to_string(),
user, user,
@@ -1226,15 +1237,15 @@ impl RagPipeline {
grounded: false, grounded: false,
refusal_reason: Some(RefusalReason::MultiHopDecomposeFailed), refusal_reason: Some(RefusalReason::MultiHopDecomposeFailed),
model: self.llm.model_ref(), model: self.llm.model_ref(),
embedding: embedding_ref_for(opts.mode, &self.config), embedding: embedding_ref_for(opts.mode, &self.models),
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(), PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(),
), ),
retrieval: AnswerRetrievalSummary { retrieval: AnswerRetrievalSummary {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: opts.k.max(self.config.search.default_k), k: opts.k.max(self.search.default_k),
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score: 0.0, top_score: 0.0,
chunks_returned: 0, chunks_returned: 0,
chunks_used: 0, chunks_used: 0,
@@ -1276,8 +1287,8 @@ impl RagPipeline {
/// (system + user) prompt to feed back into the completion budget. /// (system + user) prompt to feed back into the completion budget.
fn pack_context(&self, query: &str, hits: &[SearchHit]) -> Result<PackedContext> { fn pack_context(&self, query: &str, hits: &[SearchHit]) -> Result<PackedContext> {
// Hard ceiling for the packed-context section in tokens (≈ chars / 4). // Hard ceiling for the packed-context section in tokens (≈ chars / 4).
let cap = self.config.rag.max_context_tokens; let cap = self.rag.max_context_tokens;
let system_prompt_text = system_prompt_for(&self.config.rag.prompt_template_version)?; let system_prompt_text = system_prompt_for(&self.rag.prompt_template_version)?;
let prompt_overhead_tokens = est_tokens(system_prompt_text) + est_tokens(query) + 64; let prompt_overhead_tokens = est_tokens(system_prompt_text) + est_tokens(query) + 64;
let budget_tokens = cap.saturating_sub(prompt_overhead_tokens); let budget_tokens = cap.saturating_sub(prompt_overhead_tokens);
@@ -1369,13 +1380,13 @@ impl RagPipeline {
model: self.llm.model_ref(), model: self.llm.model_ref(),
embedding: None, embedding: None,
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
self.config.rag.prompt_template_version.clone(), self.rag.prompt_template_version.clone(),
), ),
retrieval: AnswerRetrievalSummary { retrieval: AnswerRetrievalSummary {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: k_effective, k: k_effective,
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score: 0.0, top_score: 0.0,
chunks_returned: 0, chunks_returned: 0,
chunks_used: 0, chunks_used: 0,
@@ -1421,7 +1432,7 @@ impl RagPipeline {
hops: Option<Vec<HopRecord>>, hops: Option<Vec<HopRecord>>,
) -> Result<Answer> { ) -> Result<Answer> {
let top_score = hits[0].retrieval.fusion_score; let top_score = hits[0].retrieval.fusion_score;
let gate = self.config.rag.score_gate; let gate = self.rag.score_gate;
let mut text = String::new(); let mut text = String::new();
text.push_str("근거 부족. KB에 해당 내용 없음.\n"); text.push_str("근거 부족. KB에 해당 내용 없음.\n");
text.push_str(&format!("가까운 후보 (모두 임계 {gate:.2} 미만):\n")); text.push_str(&format!("가까운 후보 (모두 임계 {gate:.2} 미만):\n"));
@@ -1461,9 +1472,9 @@ impl RagPipeline {
// semantically correct: "this answer used vector retrieval // semantically correct: "this answer used vector retrieval
// shape, even though it refused". A future reader: do not // shape, even though it refused". A future reader: do not
// "fix" this to `None`. // "fix" this to `None`.
embedding: embedding_ref_for(opts.mode, &self.config), embedding: embedding_ref_for(opts.mode, &self.models),
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
self.config.rag.prompt_template_version.clone(), self.rag.prompt_template_version.clone(),
), ),
retrieval: AnswerRetrievalSummary { retrieval: AnswerRetrievalSummary {
trace_id, trace_id,
@@ -1508,7 +1519,7 @@ impl RagPipeline {
) -> Result<Answer> { ) -> Result<Answer> {
let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX); let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX);
let trace_id = mint_trace_id(query, 0.0, &self.llm.model_ref().id); let trace_id = mint_trace_id(query, 0.0, &self.llm.model_ref().id);
let k_effective = opts.k.max(self.config.search.default_k); let k_effective = opts.k.max(self.search.default_k);
let answer = Answer { let answer = Answer {
answer: "근거 부족. 생성된 답변이 검색된 문서 내용에 충분히 entail 되지 않음." answer: "근거 부족. 생성된 답변이 검색된 문서 내용에 충분히 entail 되지 않음."
.to_string(), .to_string(),
@@ -1516,7 +1527,7 @@ impl RagPipeline {
grounded: false, grounded: false,
refusal_reason: Some(RefusalReason::NliVerificationFailed), refusal_reason: Some(RefusalReason::NliVerificationFailed),
model: self.llm.model_ref(), model: self.llm.model_ref(),
embedding: embedding_ref_for(opts.mode, &self.config), embedding: embedding_ref_for(opts.mode, &self.models),
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(), PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(),
), ),
@@ -1524,7 +1535,7 @@ impl RagPipeline {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: k_effective, k: k_effective,
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score: 0.0, top_score: 0.0,
chunks_returned: 0, chunks_returned: 0,
chunks_used: 0, chunks_used: 0,
@@ -1575,7 +1586,7 @@ impl RagPipeline {
) -> Result<Answer> { ) -> Result<Answer> {
let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX); let elapsed_ms = u32::try_from(started.elapsed().as_millis()).unwrap_or(u32::MAX);
let trace_id = mint_trace_id(query, 0.0, &self.llm.model_ref().id); let trace_id = mint_trace_id(query, 0.0, &self.llm.model_ref().id);
let k_effective = opts.k.max(self.config.search.default_k); let k_effective = opts.k.max(self.search.default_k);
let answer = Answer { let answer = Answer {
answer: "근거 부족. NLI 검증 모델을 사용할 수 없음 — `[rag] nli_threshold = 0` 으로 비활성화 후 재시도 가능." answer: "근거 부족. NLI 검증 모델을 사용할 수 없음 — `[rag] nli_threshold = 0` 으로 비활성화 후 재시도 가능."
.to_string(), .to_string(),
@@ -1583,7 +1594,7 @@ impl RagPipeline {
grounded: false, grounded: false,
refusal_reason: Some(RefusalReason::NliModelUnavailable), refusal_reason: Some(RefusalReason::NliModelUnavailable),
model: self.llm.model_ref(), model: self.llm.model_ref(),
embedding: embedding_ref_for(opts.mode, &self.config), embedding: embedding_ref_for(opts.mode, &self.models),
prompt_template_version: PromptTemplateVersion( prompt_template_version: PromptTemplateVersion(
PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(), PROMPT_TEMPLATE_VERSION_MULTI_HOP.to_string(),
), ),
@@ -1591,7 +1602,7 @@ impl RagPipeline {
trace_id, trace_id,
mode: opts.mode, mode: opts.mode,
k: k_effective, k: k_effective,
score_gate: self.config.rag.score_gate, score_gate: self.rag.score_gate,
top_score: 0.0, top_score: 0.0,
chunks_returned: 0, chunks_returned: 0,
chunks_used: 0, chunks_used: 0,
@@ -1630,13 +1641,13 @@ impl RagPipeline {
/// paths attach the configured embedding model so `kb explain` can /// paths attach the configured embedding model so `kb explain` can
/// later identify which embedder shaped the retrieval (even on /// later identify which embedder shaped the retrieval (even on
/// refusals — see `refuse_score_gate`). /// refusals — see `refuse_score_gate`).
fn embedding_ref_for(mode: SearchMode, cfg: &kebab_config::Config) -> Option<ModelRef> { fn embedding_ref_for(mode: SearchMode, models: &kebab_config::ModelsCfg) -> Option<ModelRef> {
match mode { match mode {
SearchMode::Lexical => None, SearchMode::Lexical => None,
SearchMode::Vector | SearchMode::Hybrid => Some(ModelRef { SearchMode::Vector | SearchMode::Hybrid => Some(ModelRef {
id: cfg.models.embedding.model.clone(), id: models.embedding.model.clone(),
provider: cfg.models.embedding.provider.clone(), provider: models.embedding.provider.clone(),
dimensions: Some(cfg.models.embedding.dimensions), dimensions: Some(models.embedding.dimensions),
}), }),
} }
} }

View File

@@ -37,7 +37,7 @@ impl RagEnv {
let temp = tempfile::tempdir().expect("tempdir"); let temp = tempfile::tempdir().expect("tempdir");
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = temp.path().to_string_lossy().into_owned(); config.storage.data_dir = temp.path().to_string_lossy().into_owned();
let sqlite = SqliteStore::open(&config).unwrap(); let sqlite = SqliteStore::open(&config.storage).unwrap();
sqlite.run_migrations().unwrap(); sqlite.run_migrations().unwrap();
Self { Self {
temp, temp,

View File

@@ -71,7 +71,7 @@ fn multi_hop_decide_stop_triggers_synthesize() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -138,7 +138,7 @@ fn multi_hop_decide_continue_adds_more_chunks() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -211,7 +211,7 @@ fn multi_hop_max_depth_force_stops() {
let lm = Arc::new(ScriptedLm::new(vec![r#"["q1"]"#, "answer [#1]"])); let lm = Arc::new(ScriptedLm::new(vec![r#"["q1"]"#, "answer [#1]"]));
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone());
let answer = pipeline.ask("q", multi_hop_opts()).unwrap(); let answer = pipeline.ask("q", multi_hop_opts()).unwrap();
@@ -271,7 +271,7 @@ fn multi_hop_pool_chunks_dedup_by_chunk_id() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -327,7 +327,7 @@ fn multi_hop_decide_parse_failure_falls_through_to_synthesize() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -402,7 +402,7 @@ fn multi_hop_refuse_no_chunks_preserves_hops_trace() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -492,7 +492,7 @@ fn multi_hop_refuse_score_gate_preserves_hops_trace() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -565,7 +565,7 @@ fn multi_hop_below_probe_gate_refuses_before_any_llm_call() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -608,7 +608,7 @@ fn multi_hop_empty_probe_pool_refuses_before_any_llm_call() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -653,7 +653,7 @@ fn multi_hop_above_probe_gate_proceeds_to_decompose() {
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever_dyn, retriever_dyn,
lm_dyn, lm_dyn,
env.sqlite.clone(), env.sqlite.clone(),
@@ -723,7 +723,7 @@ fn multi_hop_nli_pass_keeps_grounded() {
let verifier = MockNliVerifier::pass(); let verifier = MockNliVerifier::pass();
let verifier_handle = verifier.clone(); let verifier_handle = verifier.clone();
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();
@@ -754,7 +754,7 @@ fn multi_hop_nli_fail_refuses() {
let verifier = MockNliVerifier::fail(); let verifier = MockNliVerifier::fail();
let verifier_handle = verifier.clone(); let verifier_handle = verifier.clone();
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();
@@ -787,7 +787,7 @@ fn multi_hop_nli_disabled_skip_verify() {
let retriever_dyn: Arc<dyn Retriever> = retriever; let retriever_dyn: Arc<dyn Retriever> = retriever;
let lm_dyn: Arc<dyn LanguageModel> = lm; let lm_dyn: Arc<dyn LanguageModel> = lm;
// No `with_verifier` call — pipeline.verifier stays None. // No `with_verifier` call — pipeline.verifier stays None.
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone());
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();
@@ -810,7 +810,7 @@ fn multi_hop_nli_model_unavailable_refuses() {
let verifier = MockNliVerifier::err(); let verifier = MockNliVerifier::err();
let verifier_handle = verifier.clone(); let verifier_handle = verifier.clone();
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();

View File

@@ -65,7 +65,7 @@ fn setup_happy_pipeline_no_verifier(nli_threshold: f32) -> (RagPipeline, RagEnv)
cfg.rag.nli_threshold = nli_threshold; cfg.rag.nli_threshold = nli_threshold;
// Intentionally NO `.with_verifier()` — this is the condition under test. // Intentionally NO `.with_verifier()` — this is the condition under test.
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone());
(pipeline, env) (pipeline, env)
} }

View File

@@ -84,7 +84,7 @@ fn nli_verification_fail_emits_final_stream_event_with_refusal() {
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let (tx, rx) = mpsc::channel::<StreamEvent>(); let (tx, rx) = mpsc::channel::<StreamEvent>();
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline let answer = pipeline
@@ -134,7 +134,7 @@ fn nli_model_unavailable_emits_final_stream_event_with_refusal() {
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let (tx, rx) = mpsc::channel::<StreamEvent>(); let (tx, rx) = mpsc::channel::<StreamEvent>();
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline let answer = pipeline

View File

@@ -83,7 +83,7 @@ fn long_en_synth_answer_truncated_before_nli_call() {
let verifier_handle = verifier.clone(); let verifier_handle = verifier.clone();
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();
@@ -163,7 +163,7 @@ fn long_kr_synth_answer_retries_with_smaller_budget() {
let verifier_handle = verifier.clone(); let verifier_handle = verifier.clone();
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();
@@ -217,7 +217,7 @@ fn unrelenting_token_overflow_falls_through_to_unavailable() {
); );
let verifier_dyn: Arc<dyn NliVerifier> = verifier; let verifier_dyn: Arc<dyn NliVerifier> = verifier;
let pipeline = RagPipeline::new(cfg, retriever_dyn, lm_dyn, env.sqlite.clone()) let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever_dyn, lm_dyn, env.sqlite.clone())
.with_verifier(verifier_dyn); .with_verifier(verifier_dyn);
let answer = pipeline.ask("compound", multi_hop_opts()).unwrap(); let answer = pipeline.ask("compound", multi_hop_opts()).unwrap();

View File

@@ -82,7 +82,7 @@ fn empty_hits_refuses_no_chunks_without_llm_call() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(Vec::new())); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(Vec::new()));
let lm = Arc::new(CountingLm::new("(unused)")); let lm = Arc::new(CountingLm::new("(unused)"));
let lm_dyn: Arc<dyn LanguageModel> = lm.clone(); let lm_dyn: Arc<dyn LanguageModel> = lm.clone();
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm_dyn, env.sqlite.clone());
let answer = pipeline.ask("anything", default_opts()).unwrap(); let answer = pipeline.ask("anything", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::NoChunks)); assert_eq!(answer.refusal_reason, Some(RefusalReason::NoChunks));
@@ -105,7 +105,7 @@ fn top_below_gate_refuses_score_gate_without_llm_call() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm = Arc::new(CountingLm::new("(unused)")); let lm = Arc::new(CountingLm::new("(unused)"));
let lm_dyn: Arc<dyn LanguageModel> = lm.clone(); let lm_dyn: Arc<dyn LanguageModel> = lm.clone();
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm_dyn, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::ScoreGate)); assert_eq!(answer.refusal_reason, Some(RefusalReason::ScoreGate));
@@ -142,7 +142,7 @@ fn grounded_happy_path_marker_one() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let canned = "Rust is a systems language. [#1]"; let canned = "Rust is a systems language. [#1]";
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned)); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("what is rust", default_opts()).unwrap(); let answer = pipeline.ask("what is rust", default_opts()).unwrap();
assert!(answer.grounded); assert!(answer.grounded);
@@ -165,7 +165,7 @@ fn unknown_marker_refuses_llm_self_judge() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
// Marker 7 is NOT in the packed set (only #1 is). // Marker 7 is NOT in the packed set (only #1 is).
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("answer text [#7]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("answer text [#7]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge)); assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge));
@@ -187,7 +187,7 @@ fn marker_without_hash_is_no_marker() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
// `[1]` is NOT a valid marker — strict regex requires `[#1]`. // `[1]` is NOT a valid marker — strict regex requires `[#1]`.
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("the answer [1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("the answer [1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge)); assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge));
@@ -206,7 +206,7 @@ fn vec_bracket_one_is_no_false_positive() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
// `vec![1]` MUST NOT be misread as a citation marker. // `vec![1]` MUST NOT be misread as a citation marker.
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("see vec![1] in code")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("see vec![1] in code"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge)); assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge));
@@ -224,7 +224,7 @@ fn explicit_korean_refusal_is_self_judge() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("근거가 부족합니다.")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("근거가 부족합니다."));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge)); assert_eq!(answer.refusal_reason, Some(RefusalReason::LlmSelfJudge));
@@ -263,7 +263,7 @@ fn packing_stops_before_budget_overflow() {
} }
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("ok [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("ok [#1]"));
let pipeline = RagPipeline::new(cfg, retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(cfg.rag.clone(), cfg.models.clone(), cfg.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
// At least one chunk was packed; the budget cap should keep it to <= 1. // At least one chunk was packed; the budget cap should keep it to <= 1.
@@ -287,7 +287,7 @@ fn streaming_forwards_tokens_to_sink() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let canned = "ok [#1]"; let canned = "ok [#1]";
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned)); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let (tx, rx) = std::sync::mpsc::channel::<StreamEvent>(); let (tx, rx) = std::sync::mpsc::channel::<StreamEvent>();
let mut opts = default_opts(); let mut opts = default_opts();
@@ -322,7 +322,7 @@ fn dropped_receiver_aborts_with_llm_stream_aborted() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let canned = "ok [#1]"; let canned = "ok [#1]";
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned)); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let (tx, rx) = std::sync::mpsc::channel::<StreamEvent>(); let (tx, rx) = std::sync::mpsc::channel::<StreamEvent>();
drop(rx); // receiver gone — first Token send fails, loop breaks drop(rx); // receiver gone — first Token send fails, loop breaks
@@ -352,7 +352,7 @@ fn usage_populated_from_done_chunk() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("ok [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("ok [#1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.usage.prompt_tokens, 10, "from canned_usage"); assert_eq!(answer.usage.prompt_tokens, 10, "from canned_usage");
@@ -368,7 +368,7 @@ fn answers_row_inserted_for_each_refusal_kind() {
let env = RagEnv::new(); let env = RagEnv::new();
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(Vec::new())); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(Vec::new()));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(""));
let p = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let p = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
p.ask("q", default_opts()).unwrap(); p.ask("q", default_opts()).unwrap();
assert_eq!(env.count_answers(), 1); assert_eq!(env.count_answers(), 1);
} }
@@ -381,7 +381,7 @@ fn answers_row_inserted_for_each_refusal_kind() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.05, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.05, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(""));
let p = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let p = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
p.ask("q", default_opts()).unwrap(); p.ask("q", default_opts()).unwrap();
assert_eq!(env.count_answers(), 1); assert_eq!(env.count_answers(), 1);
} }
@@ -394,7 +394,7 @@ fn answers_row_inserted_for_each_refusal_kind() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("answer with no marker")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("answer with no marker"));
let p = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let p = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
p.ask("q", default_opts()).unwrap(); p.ask("q", default_opts()).unwrap();
assert_eq!(env.count_answers(), 1); assert_eq!(env.count_answers(), 1);
} }
@@ -413,7 +413,7 @@ fn determinism_temperature_zero_seed_zero() {
let mk_pipeline = || { let mk_pipeline = || {
let r: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits.clone())); let r: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits.clone()));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]"));
RagPipeline::new(env.config.clone(), r, lm, env.sqlite.clone()) RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), r, lm, env.sqlite.clone())
}; };
let a1 = mk_pipeline().ask("q", default_opts()).unwrap(); let a1 = mk_pipeline().ask("q", default_opts()).unwrap();
let a2 = mk_pipeline().ask("q", default_opts()).unwrap(); let a2 = mk_pipeline().ask("q", default_opts()).unwrap();
@@ -443,7 +443,7 @@ fn unfetchable_chunks_fall_back_to_no_chunks() {
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm = Arc::new(CountingLm::new("(should never run)")); let lm = Arc::new(CountingLm::new("(should never run)"));
let lm_dyn: Arc<dyn LanguageModel> = lm.clone(); let lm_dyn: Arc<dyn LanguageModel> = lm.clone();
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm_dyn, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm_dyn, env.sqlite.clone());
let answer = pipeline.ask("q", default_opts()).unwrap(); let answer = pipeline.ask("q", default_opts()).unwrap();
assert_eq!(answer.refusal_reason, Some(RefusalReason::NoChunks)); assert_eq!(answer.refusal_reason, Some(RefusalReason::NoChunks));
@@ -484,7 +484,7 @@ fn grounded_citations_inherit_indexed_at_and_stale_from_hit() {
)]; )];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("apples are fruit. [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("apples are fruit. [#1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("apples", default_opts()).unwrap(); let answer = pipeline.ask("apples", default_opts()).unwrap();
assert!(answer.grounded); assert!(answer.grounded);
@@ -523,7 +523,7 @@ fn grounded_citations_not_stale_for_fresh_hit() {
)]; )];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("apples are fruit. [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("apples are fruit. [#1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("apples", default_opts()).unwrap(); let answer = pipeline.ask("apples", default_opts()).unwrap();
assert!(answer.grounded); assert!(answer.grounded);
@@ -553,7 +553,7 @@ fn answer_json_serializes_with_expected_keys() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("what", default_opts()).unwrap(); let answer = pipeline.ask("what", default_opts()).unwrap();
let v: serde_json::Value = serde_json::to_value(&answer).unwrap(); let v: serde_json::Value = serde_json::to_value(&answer).unwrap();
// Stable top-level key set per `answer.v1` (§2.3). // Stable top-level key set per `answer.v1` (§2.3).
@@ -611,7 +611,7 @@ fn ask_multi_hop_dispatches_and_decompose_garbage_refuses() {
let lm = Arc::new(CountingLm::new("definitely not a JSON array")); let lm = Arc::new(CountingLm::new("definitely not a JSON array"));
let lm_handle = lm.clone(); let lm_handle = lm.clone();
let pipeline = RagPipeline::new( let pipeline = RagPipeline::new(
env.config.clone(), env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(),
retriever, retriever,
lm.clone() as Arc<dyn LanguageModel>, lm.clone() as Arc<dyn LanguageModel>,
env.sqlite.clone(), env.sqlite.clone(),
@@ -665,7 +665,7 @@ fn ask_with_multi_hop_false_keeps_single_pass_path() {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]")); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new("Rust is. [#1]"));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let answer = pipeline.ask("what", default_opts()).unwrap(); let answer = pipeline.ask("what", default_opts()).unwrap();

View File

@@ -112,7 +112,7 @@ fn build_pipeline_with_template(
env.seed_chunk(&chunk_id, &doc_id, "a.md", "hello world", &["H"]); env.seed_chunk(&chunk_id, &doc_id, "a.md", "hello world", &["H"]);
let hit = mk_hit(1, &chunk_id, &doc_id, "a.md", 0.9, &["H"]); let hit = mk_hit(1, &chunk_id, &doc_id, "a.md", 0.9, &["H"]);
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(vec![hit])); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(vec![hit]));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
(pipeline, captured, env) (pipeline, captured, env)
} }
@@ -199,7 +199,7 @@ fn pack_user_prompt_for_hit(
hit.source_id = source_id.map(str::to_string); hit.source_id = source_id.map(str::to_string);
hit.trust_level = trust_level; hit.trust_level = trust_level;
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(vec![hit])); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(vec![hit]));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let _ = pipeline.ask("hello", lexical_opts()); let _ = pipeline.ask("hello", lexical_opts());
let out = captured_user let out = captured_user
.lock() .lock()

View File

@@ -81,7 +81,7 @@ fn env_with_one_hit(canned: &str) -> (RagEnv, RagPipeline) {
let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])]; let hits = vec![mk_hit(1, &cid, &did, "notes/a.md", 0.85, &["Intro"])];
let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits)); let retriever: Arc<dyn Retriever> = Arc::new(MockRetriever::new(hits));
let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned)); let lm: Arc<dyn LanguageModel> = Arc::new(CountingLm::new(canned));
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
(env, pipeline) (env, pipeline)
} }
@@ -186,7 +186,7 @@ fn ask_emits_no_final_when_cancelled_mid_stream() {
}, },
gate: Arc::clone(&gate), gate: Arc::clone(&gate),
}); });
let pipeline = RagPipeline::new(env.config.clone(), retriever, lm, env.sqlite.clone()); let pipeline = RagPipeline::new(env.config.rag.clone(), env.config.models.clone(), env.config.search.clone(), retriever, lm, env.sqlite.clone());
let (tx, rx) = mpsc::channel::<StreamEvent>(); let (tx, rx) = mpsc::channel::<StreamEvent>();
let opts = opts_with_sink(tx); let opts = opts_with_sink(tx);

View File

@@ -74,19 +74,19 @@ pub struct HybridRetriever {
} }
impl HybridRetriever { impl HybridRetriever {
/// Construct from a `kb-config` Config + the two underlying /// Construct from the `[search]` config slice + the two underlying
/// retrievers. Reads `config.search.hybrid_fusion` (only `"rrf"` /// retrievers. Reads `search.hybrid_fusion` (only `"rrf"`
/// is recognised today) and `config.search.rrf_k`. /// is recognised today) and `search.rrf_k`.
pub fn new( pub fn new(
config: &kebab_config::Config, search: &kebab_config::SearchCfg,
lexical: Arc<dyn Retriever>, lexical: Arc<dyn Retriever>,
vector: Arc<dyn Retriever>, vector: Arc<dyn Retriever>,
) -> Self { ) -> Self {
let fusion = parse_fusion(&config.search.hybrid_fusion, config.search.rrf_k); let fusion = parse_fusion(&search.hybrid_fusion, search.rrf_k);
let default_k = if config.search.default_k == 0 { let default_k = if search.default_k == 0 {
DEFAULT_K DEFAULT_K
} else { } else {
config.search.default_k search.default_k
}; };
// Surface mismatched index_version up front so users see it // Surface mismatched index_version up front so users see it
// (e.g. lexical at v2, vector at v1 means a stale index that // (e.g. lexical at v2, vector at v1 means a stale index that

View File

@@ -61,24 +61,18 @@ pub struct VectorRetriever {
impl VectorRetriever { impl VectorRetriever {
/// Construct with `index_version` derived from the configured /// Construct with `index_version` derived from the configured
/// embedding model + dimensions, and snippet width pulled from /// embedding model + dimensions and an explicit `snippet_chars`
/// `kb-config`'s defaults. /// (the caller passes `config.search.snippet_chars`).
/// ///
/// The explicit `index_version` form is [`Self::with_settings`]. /// Thin delegate to [`Self::with_settings`].
pub fn new( pub fn new(
store: Arc<dyn VectorStore + Send + Sync>, store: Arc<dyn VectorStore + Send + Sync>,
embed: Arc<dyn Embedder>, embed: Arc<dyn Embedder>,
sqlite: Arc<SqliteStore>, sqlite: Arc<SqliteStore>,
index_version: IndexVersion, index_version: IndexVersion,
snippet_chars: usize,
) -> Self { ) -> Self {
let cfg = kebab_config::Config::defaults(); Self::with_settings(store, embed, sqlite, index_version, snippet_chars)
Self::with_settings(
store,
embed,
sqlite,
index_version,
cfg.search.snippet_chars,
)
} }
/// Construct with explicit `snippet_chars`. Mirrors the lexical /// Construct with explicit `snippet_chars`. Mirrors the lexical

View File

@@ -68,10 +68,10 @@ impl HybridEnv {
let temp = tempfile::tempdir().expect("tempdir"); let temp = tempfile::tempdir().expect("tempdir");
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = temp.path().to_string_lossy().into_owned(); config.storage.data_dir = temp.path().to_string_lossy().into_owned();
let sqlite = SqliteStore::open(&config).unwrap(); let sqlite = SqliteStore::open(&config.storage).unwrap();
sqlite.run_migrations().unwrap(); sqlite.run_migrations().unwrap();
let sqlite = Arc::new(sqlite); let sqlite = Arc::new(sqlite);
let vector_store = Arc::new(LanceVectorStore::new(&config, sqlite.clone()).unwrap()); let vector_store = Arc::new(LanceVectorStore::new(&config.storage, sqlite.clone()).unwrap());
let embedder = Arc::new(MockEmbedder::new( let embedder = Arc::new(MockEmbedder::new(
EmbeddingModelId(TEST_MODEL_ID.to_string()), EmbeddingModelId(TEST_MODEL_ID.to_string()),
EmbeddingVersion("v1".to_string()), EmbeddingVersion("v1".to_string()),
@@ -105,6 +105,7 @@ impl HybridEnv {
embed, embed,
Arc::clone(&self.sqlite), Arc::clone(&self.sqlite),
IndexVersion(TEST_VEC_INDEX_VERSION.to_string()), IndexVersion(TEST_VEC_INDEX_VERSION.to_string()),
self.config.search.snippet_chars,
) )
} }

View File

@@ -31,7 +31,7 @@ impl Env {
let temp = tempfile::tempdir().expect("tempdir"); let temp = tempfile::tempdir().expect("tempdir");
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = temp.path().to_string_lossy().into_owned(); config.storage.data_dir = temp.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&config).expect("open store"); let store = SqliteStore::open(&config.storage).expect("open store");
store.run_migrations().expect("run migrations"); store.run_migrations().expect("run migrations");
let db_path = temp.path().join("kebab.sqlite"); let db_path = temp.path().join("kebab.sqlite");
Self { Self {

View File

@@ -118,7 +118,7 @@ mod tests {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
(dir, store) (dir, store)
} }

View File

@@ -213,7 +213,7 @@ mod tests {
fn open_store(tmp: &TempDir) -> SqliteStore { fn open_store(tmp: &TempDir) -> SqliteStore {
let cfg = config_for(tmp); let cfg = config_for(tmp);
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
} }

View File

@@ -310,7 +310,7 @@ mod tests {
fn open_store(tmp: &TempDir) -> SqliteStore { fn open_store(tmp: &TempDir) -> SqliteStore {
let mut c = Config::defaults(); let mut c = Config::defaults();
c.storage.data_dir = tmp.path().to_string_lossy().into_owned(); c.storage.data_dir = tmp.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&c).unwrap(); let store = SqliteStore::open(&c.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
} }

View File

@@ -117,7 +117,7 @@ mod tests {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
let store = crate::SqliteStore::open(&cfg).unwrap(); let store = crate::SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
(dir, store) (dir, store)
} }

View File

@@ -116,12 +116,12 @@ impl SqliteStore {
}) })
} }
/// Open (or create) the SQLite file under `config.storage.data_dir`, /// Open (or create) the SQLite file under `storage.data_dir`,
/// apply pragmas (foreign_keys / WAL / synchronous=NORMAL / /// apply pragmas (foreign_keys / WAL / synchronous=NORMAL /
/// temp_store=MEMORY), and create parent directories as needed. /// temp_store=MEMORY), and create parent directories as needed.
/// **Does not run migrations** — call [`Self::run_migrations`] next. /// **Does not run migrations** — call [`Self::run_migrations`] next.
pub fn open(config: &kebab_config::Config) -> Result<Self> { pub fn open(storage: &kebab_config::StorageCfg) -> Result<Self> {
let data_dir = kebab_config::expand_path(&config.storage.data_dir, ""); let data_dir = kebab_config::expand_path(&storage.data_dir, "");
std::fs::create_dir_all(&data_dir) std::fs::create_dir_all(&data_dir)
.with_context(|| format!("create data_dir {}", data_dir.display()))?; .with_context(|| format!("create data_dir {}", data_dir.display()))?;
let db_path = data_dir.join(SQLITE_FILE); let db_path = data_dir.join(SQLITE_FILE);
@@ -139,7 +139,7 @@ impl SqliteStore {
Ok(Self { Ok(Self {
data_dir, data_dir,
copy_threshold_bytes: config.storage.copy_threshold_mb * BYTES_PER_MIB, copy_threshold_bytes: storage.copy_threshold_mb * BYTES_PER_MIB,
conn: Mutex::new(conn), conn: Mutex::new(conn),
}) })
} }
@@ -1189,7 +1189,7 @@ mod tests {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
let mut cfg = kebab_config::Config::defaults(); let mut cfg = kebab_config::Config::defaults();
cfg.storage.data_dir = dir.path().to_string_lossy().into_owned(); cfg.storage.data_dir = dir.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
(dir, store) (dir, store)
} }

View File

@@ -33,7 +33,7 @@ fn b3_full_hex(bytes: &[u8]) -> String {
#[test] #[test]
fn copy_mode_writes_file_with_0o644_and_correct_bytes() { fn copy_mode_writes_file_with_0o644_and_correct_bytes() {
let env = common::TestEnv::with_threshold(100); let env = common::TestEnv::with_threshold(100);
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let bytes = b"hello, sqlite"; let bytes = b"hello, sqlite";
@@ -80,7 +80,7 @@ fn copy_mode_writes_file_with_0o644_and_correct_bytes() {
fn reference_mode_does_not_write_file_but_records_path() { fn reference_mode_does_not_write_file_but_records_path() {
// copy_threshold_mb=0 → every byte lands on the reference branch. // copy_threshold_mb=0 → every byte lands on the reference branch.
let env = common::TestEnv::with_threshold(0); let env = common::TestEnv::with_threshold(0);
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let bytes = b"big-pretend-bytes"; let bytes = b"big-pretend-bytes";
@@ -126,7 +126,7 @@ fn put_asset_with_bytes_sweeps_workspace_path_orphan() {
// is exercised end-to-end in `kebab-app::tests::pdf_pipeline:: // is exercised end-to-end in `kebab-app::tests::pdf_pipeline::
// re_ingest_edited_pdf_produces_new_doc_id`. // re_ingest_edited_pdf_produces_new_doc_id`.
let env = common::TestEnv::with_threshold(100); let env = common::TestEnv::with_threshold(100);
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// Pre-populate a row that owns `notes/foo.md` under a *different* // Pre-populate a row that owns `notes/foo.md` under a *different*
@@ -200,7 +200,7 @@ fn put_asset_with_bytes_rejects_invalid_asset_id() {
// 32-hex `FromStr` invariant. The store boundary must reject any ID // 32-hex `FromStr` invariant. The store boundary must reject any ID
// whose shape would let path construction escape `data_dir/assets/`. // whose shape would let path construction escape `data_dir/assets/`.
let env = common::TestEnv::with_threshold(100); let env = common::TestEnv::with_threshold(100);
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// 32 chars but contains a `/` — would let `assets_path_for` stitch // 32 chars but contains a `/` — would let `assets_path_for` stitch
@@ -242,7 +242,7 @@ fn put_asset_with_bytes_rejects_invalid_asset_id() {
#[test] #[test]
fn checksum_mismatch_returns_conflict() { fn checksum_mismatch_returns_conflict() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let bytes = b"the real bytes"; let bytes = b"the real bytes";

View File

@@ -30,7 +30,7 @@ fn fixtures_dir() -> PathBuf {
#[test] #[test]
fn document_and_chunks_round_trip_through_sqlite() { fn document_and_chunks_round_trip_through_sqlite() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// ── Build inputs from the fixture ─────────────────────────────── // ── Build inputs from the fixture ───────────────────────────────

View File

@@ -15,7 +15,7 @@ fn config_for(tmp: &TempDir) -> Config {
fn open_store(tmp: &TempDir) -> SqliteStore { fn open_store(tmp: &TempDir) -> SqliteStore {
let cfg = config_for(tmp); let cfg = config_for(tmp);
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
} }

View File

@@ -18,7 +18,7 @@ use time::OffsetDateTime;
fn open_store(tmp: &TempDir) -> SqliteStore { fn open_store(tmp: &TempDir) -> SqliteStore {
let mut c = Config::defaults(); let mut c = Config::defaults();
c.storage.data_dir = tmp.path().to_string_lossy().into_owned(); c.storage.data_dir = tmp.path().to_string_lossy().into_owned();
let store = SqliteStore::open(&c).unwrap(); let store = SqliteStore::open(&c.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
} }

View File

@@ -133,7 +133,7 @@ fn fts_v002_backfills_existing_chunks() {
#[test] #[test]
fn fts_v002_backfill_select_matches_chunks_count() { fn fts_v002_backfill_select_matches_chunks_count() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -158,7 +158,7 @@ fn fts_v002_backfill_select_matches_chunks_count() {
#[test] #[test]
fn fts_chunks_ai_trigger_propagates_insert() { fn fts_chunks_ai_trigger_propagates_insert() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -185,7 +185,7 @@ fn fts_chunks_ai_trigger_propagates_insert() {
#[test] #[test]
fn fts_chunks_ad_trigger_propagates_delete() { fn fts_chunks_ad_trigger_propagates_delete() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -205,7 +205,7 @@ fn fts_chunks_ad_trigger_propagates_delete() {
#[test] #[test]
fn fts_chunks_au_trigger_propagates_update() { fn fts_chunks_au_trigger_propagates_update() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -246,7 +246,7 @@ fn count_match(conn: &Connection, term: &str) -> i64 {
#[test] #[test]
fn fts_rebuild_chunks_fts_is_idempotent() { fn fts_rebuild_chunks_fts_is_idempotent() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -274,7 +274,7 @@ fn fts_rebuild_chunks_fts_is_idempotent() {
#[test] #[test]
fn fts_rebuild_chunks_fts_recovers_from_drift() { fn fts_rebuild_chunks_fts_recovers_from_drift() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -297,7 +297,7 @@ fn fts_rebuild_chunks_fts_recovers_from_drift() {
#[test] #[test]
fn fts_double_run_migrations_is_noop() { fn fts_double_run_migrations_is_noop() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().expect("run 1"); store.run_migrations().expect("run 1");
// Second invocation must be a no-op (refinery's bookkeeping table // Second invocation must be a no-op (refinery's bookkeeping table
// tracks applied versions). The chunks_fts virtual table is still // tracks applied versions). The chunks_fts virtual table is still
@@ -444,7 +444,7 @@ fn fts_v009_matches_design_section_5_5_verbatim() {
#[test] #[test]
fn v009_bumps_corpus_revision() { fn v009_bumps_corpus_revision() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let rev = store.corpus_revision(); let rev = store.corpus_revision();
assert!( assert!(
@@ -459,7 +459,7 @@ fn v009_bumps_corpus_revision() {
#[test] #[test]
fn backfill_tokenized_korean_text_populates_nullable_rows() { fn backfill_tokenized_korean_text_populates_nullable_rows() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// chunks 에 한국어 row 두 개 INSERT (tokenized_korean_text 는 chunks_ai trigger // chunks 에 한국어 row 두 개 INSERT (tokenized_korean_text 는 chunks_ai trigger
@@ -527,7 +527,7 @@ fn fts_store_drop_releases_wal_files() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let db_path = env.db_path(); let db_path = env.db_path();
{ {
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// Force at least one trigger fire so WAL has content to flush. // Force at least one trigger fire so WAL has content to flush.
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -575,7 +575,7 @@ fn fts_store_drop_releases_wal_files() {
#[test] #[test]
fn fts_v009_unicode61_space_separated_korean_token_hits() { fn fts_v009_unicode61_space_separated_korean_token_hits() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -605,7 +605,7 @@ fn fts_v009_unicode61_space_separated_korean_token_hits() {
#[test] #[test]
fn fts_v009_korean_morphological_2char_query_hits() { fn fts_v009_korean_morphological_2char_query_hits() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);
@@ -633,7 +633,7 @@ fn fts_v009_korean_morphological_2char_query_hits() {
#[test] #[test]
fn fts_v009_english_whole_token_only() { fn fts_v009_english_whole_token_only() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let conn = raw_conn_no_fk(&env); let conn = raw_conn_no_fk(&env);

View File

@@ -105,7 +105,7 @@ fn make_chunks(doc_id: &DocumentId) -> Vec<Chunk> {
#[test] #[test]
fn put_document_idempotent_bumps_doc_version() { fn put_document_idempotent_bumps_doc_version() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();
@@ -149,7 +149,7 @@ fn put_document_idempotent_bumps_doc_version() {
#[test] #[test]
fn put_blocks_and_put_chunks_replace_not_duplicate() { fn put_blocks_and_put_chunks_replace_not_duplicate() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();
@@ -209,7 +209,7 @@ fn put_blocks_and_put_chunks_replace_not_duplicate() {
#[test] #[test]
fn put_blocks_transactional_rollback_on_fk_violation() { fn put_blocks_transactional_rollback_on_fk_violation() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();

View File

@@ -77,7 +77,7 @@ fn make_doc() -> CanonicalDocument {
#[test] #[test]
fn put_then_get_document_roundtrips_version_stamps() { fn put_then_get_document_roundtrips_version_stamps() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();
@@ -100,7 +100,7 @@ fn put_then_get_document_roundtrips_version_stamps() {
#[test] #[test]
fn put_then_get_document_roundtrips_none_stamps() { fn put_then_get_document_roundtrips_none_stamps() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();
@@ -126,7 +126,7 @@ fn put_then_get_document_roundtrips_none_stamps() {
#[test] #[test]
fn get_asset_by_workspace_path_roundtrips() { fn get_asset_by_workspace_path_roundtrips() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let asset = make_asset(); let asset = make_asset();
@@ -145,7 +145,7 @@ fn get_asset_by_workspace_path_roundtrips() {
#[test] #[test]
fn get_asset_by_workspace_path_returns_none_for_unknown() { fn get_asset_by_workspace_path_returns_none_for_unknown() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let path = WorkspacePath::new("notes/missing.md".into()).unwrap(); let path = WorkspacePath::new("notes/missing.md".into()).unwrap();

View File

@@ -9,7 +9,7 @@ mod common;
#[test] #[test]
fn create_then_progress_then_finish() { fn create_then_progress_then_finish() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let id = store let id = store
@@ -39,7 +39,7 @@ fn create_then_progress_then_finish() {
#[test] #[test]
fn finish_with_error_message_is_round_trippable() { fn finish_with_error_message_is_round_trippable() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
let id = store.create(JobKind::Embed, json!({})).unwrap(); let id = store.create(JobKind::Embed, json!({})).unwrap();
@@ -59,7 +59,7 @@ fn finish_with_error_message_is_round_trippable() {
#[test] #[test]
fn list_filters_status_and_kind() { fn list_filters_status_and_kind() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
// Two ingest jobs (one finished succeeded, one pending) + one embed. // Two ingest jobs (one finished succeeded, one pending) + one embed.

View File

@@ -81,7 +81,7 @@ fn make_doc(
#[test] #[test]
fn list_documents_filters_lang_and_tags() { fn list_documents_filters_lang_and_tags() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).unwrap(); let store = SqliteStore::open(&env.config().storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
for (asset, doc) in [ for (asset, doc) in [

View File

@@ -8,7 +8,7 @@ mod common;
#[test] #[test]
fn fresh_db_has_all_p1_tables_and_indexes() { fn fresh_db_has_all_p1_tables_and_indexes() {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).expect("open"); let store = SqliteStore::open(&env.config().storage).expect("open");
store.run_migrations().expect("run migrations"); store.run_migrations().expect("run migrations");
// Pull the list of user tables from sqlite_master. // Pull the list of user tables from sqlite_master.

View File

@@ -8,7 +8,7 @@ use rusqlite::OptionalExtension;
fn open_migrated() -> (common::TestEnv, SqliteStore) { fn open_migrated() -> (common::TestEnv, SqliteStore) {
let env = common::TestEnv::new(); let env = common::TestEnv::new();
let store = SqliteStore::open(&env.config()).expect("open"); let store = SqliteStore::open(&env.config().storage).expect("open");
store.run_migrations().expect("run migrations"); store.run_migrations().expect("run migrations");
(env, store) (env, store)
} }

View File

@@ -19,7 +19,7 @@ fn config_for(tmp: &TempDir) -> Config {
fn open_store(tmp: &TempDir) -> SqliteStore { fn open_store(tmp: &TempDir) -> SqliteStore {
let cfg = config_for(tmp); let cfg = config_for(tmp);
let store = SqliteStore::open(&cfg).unwrap(); let store = SqliteStore::open(&cfg.storage).unwrap();
store.run_migrations().unwrap(); store.run_migrations().unwrap();
store store
} }

View File

@@ -83,7 +83,7 @@ pub struct LanceVectorStore {
impl LanceVectorStore { impl LanceVectorStore {
/// Open (or create) the Lance directory under /// Open (or create) the Lance directory under
/// `config.storage.vector_dir`, build a current-thread tokio /// `storage.vector_dir`, build a current-thread tokio
/// runtime, and return a ready-to-use store. Migrations on the /// runtime, and return a ready-to-use store. Migrations on the
/// SQLite side must already have been applied (`run_migrations`) /// SQLite side must already have been applied (`run_migrations`)
/// — this constructor does not touch the SQLite schema. /// — this constructor does not touch the SQLite schema.
@@ -93,9 +93,9 @@ impl LanceVectorStore {
/// runtime context will panic with `"Cannot start a runtime from /// runtime context will panic with `"Cannot start a runtime from
/// within a runtime"`. See the struct-level `# Async context` /// within a runtime"`. See the struct-level `# Async context`
/// section. /// section.
pub fn new(config: &kebab_config::Config, sqlite: Arc<SqliteStore>) -> Result<Self> { pub fn new(storage: &kebab_config::StorageCfg, sqlite: Arc<SqliteStore>) -> Result<Self> {
let data_dir = expand_path(&config.storage.data_dir, ""); let data_dir = expand_path(&storage.data_dir, "");
let vector_dir = expand_path(&config.storage.vector_dir, &data_dir.to_string_lossy()); let vector_dir = expand_path(&storage.vector_dir, &data_dir.to_string_lossy());
std::fs::create_dir_all(&vector_dir) std::fs::create_dir_all(&vector_dir)
.with_context(|| format!("create vector_dir {}", vector_dir.display()))?; .with_context(|| format!("create vector_dir {}", vector_dir.display()))?;

View File

@@ -79,10 +79,10 @@ impl TestEnv {
let temp = tempfile::tempdir().expect("tempdir"); let temp = tempfile::tempdir().expect("tempdir");
let mut config = Config::defaults(); let mut config = Config::defaults();
config.storage.data_dir = temp.path().to_string_lossy().into_owned(); config.storage.data_dir = temp.path().to_string_lossy().into_owned();
let sqlite = SqliteStore::open(&config).unwrap(); let sqlite = SqliteStore::open(&config.storage).unwrap();
sqlite.run_migrations().unwrap(); sqlite.run_migrations().unwrap();
let sqlite = Arc::new(sqlite); let sqlite = Arc::new(sqlite);
let vector = LanceVectorStore::new(&config, sqlite.clone()).unwrap(); let vector = LanceVectorStore::new(&config.storage, sqlite.clone()).unwrap();
Self { Self {
temp, temp,
config, config,