refactor(config): consumer들이 &Config 대신 타입 슬라이스 수령 (god-struct 결합 해소)
This commit is contained in:
@@ -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();
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|
||||||
|
|||||||
@@ -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 {})",
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -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 \
|
||||||
|
|||||||
@@ -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 })
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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");
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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")?;
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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),
|
||||||
}),
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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";
|
||||||
|
|||||||
@@ -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 ───────────────────────────────
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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 [
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()))?;
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user