1use std::collections::BTreeMap;
2use std::fs;
3use std::path::{Path, PathBuf};
4use std::time::Instant;
5
6use anyhow::{anyhow, bail, Context, Result};
7use rusqlite::Connection;
8use serde::Serialize;
9use toml_edit::{value, DocumentMut, Item, Table};
10
11use super::golden::{self, CategoryEvaluation, GoldenDataset};
12use crate::retrieval::embedding::{
13 self, EmbeddingConfig, EmbeddingProvider, EmbeddingProviderStatus,
14};
15
16pub const DEFAULT_DATASET_PATH: &str = "eval/golden.json";
17pub const DEFAULT_REPORT_PATH: &str = "eval/provider-comparison/report.json";
18
19const REPORT_VERSION: &str = "2026-07-29";
20const PROVIDER_COMPARISON_SLICE: &str = "provider_comparison";
21const QUERY_EMBEDDING_LATENCY_BUDGET_P95_MS: f64 = 1000.0;
22const COLD_START_EMBEDDING_LATENCY_BUDGET_MS: f64 = 1000.0;
23const EXISTING_REGRESSION_BUDGET: f64 = 0.0;
24const EPSILON: f64 = 0.000_001;
25
26mod decision;
27mod display;
28mod model_pin;
29
30use decision::build_default_decision;
31
32const ENV_CONFIG: &str = "REMEM_CONFIG";
33const ENV_PROVIDER: &str = "REMEM_EMBEDDINGS_PROVIDER";
34const ENV_PROVIDER_LEGACY: &str = "REMEM_EMBEDDING_PROVIDER";
35const ENV_MODEL: &str = "REMEM_EMBEDDINGS_MODEL";
36const ENV_MODEL_LEGACY: &str = "REMEM_EMBEDDING_MODEL";
37const ENV_BASE_URL: &str = "REMEM_EMBEDDINGS_BASE_URL";
38const ENV_BASE_URL_LEGACY: &str = "REMEM_EMBEDDING_BASE_URL";
39const ENV_DIMENSIONS: &str = "REMEM_EMBEDDINGS_DIMENSIONS";
40const ENV_DIMENSIONS_LEGACY: &str = "REMEM_EMBEDDING_DIMENSIONS";
41const ENV_API_KEY: &str = "REMEM_EMBEDDINGS_API_KEY";
42const ENV_API_KEY_LEGACY: &str = "REMEM_EMBEDDING_API_KEY";
43const ENV_API_KEY_ENV: &str = "REMEM_EMBEDDINGS_API_KEY_ENV";
44const ENV_TIMEOUT_SECS: &str = "REMEM_EMBEDDINGS_TIMEOUT_SECS";
45const ENV_FALLBACK: &str = "REMEM_EMBEDDINGS_FALLBACK";
46const ENV_MODEL_DIR: &str = "REMEM_EMBEDDINGS_MODEL_DIR";
47const ENV_HF_HOME: &str = "HF_HOME";
48const DEFAULT_API_KEY_ENV: &str = "OPENAI_API_KEY";
49
50const BASE_ENV_KEYS: &[&str] = &[
51 ENV_CONFIG,
52 ENV_PROVIDER,
53 ENV_PROVIDER_LEGACY,
54 ENV_MODEL,
55 ENV_MODEL_LEGACY,
56 ENV_BASE_URL,
57 ENV_BASE_URL_LEGACY,
58 ENV_DIMENSIONS,
59 ENV_DIMENSIONS_LEGACY,
60 ENV_API_KEY,
61 ENV_API_KEY_LEGACY,
62 ENV_API_KEY_ENV,
63 ENV_TIMEOUT_SECS,
64 ENV_FALLBACK,
65 ENV_MODEL_DIR,
66 ENV_HF_HOME,
67 DEFAULT_API_KEY_ENV,
68];
69
70#[derive(Debug, Clone)]
71pub struct ProviderComparisonOptions {
72 pub dataset_path: String,
73 pub k: usize,
74 pub json_out: String,
75 pub allow_api: bool,
76}
77
78impl Default for ProviderComparisonOptions {
79 fn default() -> Self {
80 Self {
81 dataset_path: DEFAULT_DATASET_PATH.to_string(),
82 k: 5,
83 json_out: DEFAULT_REPORT_PATH.to_string(),
84 allow_api: false,
85 }
86 }
87}
88
89#[derive(Debug, Clone, Serialize)]
90pub struct ProviderComparisonReport {
91 pub version: &'static str,
92 pub generated_at_epoch: i64,
93 pub build_profile: &'static str,
94 pub target_os: &'static str,
95 pub target_arch: &'static str,
96 pub dataset_path: String,
97 pub k: usize,
98 pub required_providers: Vec<&'static str>,
99 pub provider_comparison_slice: &'static str,
100 pub query_embedding_latency_budget_p95_ms: f64,
101 pub cold_start_embedding_latency_budget_ms: f64,
102 pub existing_regression_budget: f64,
103 pub providers: Vec<ProviderComparisonRow>,
104 pub default_decision: DefaultDecision,
105 pub notes: Vec<&'static str>,
106}
107
108#[derive(Debug, Clone, Serialize)]
109pub struct ProviderComparisonRow {
110 pub provider: &'static str,
111 pub configured_provider: String,
112 pub active_provider: String,
113 pub fallback_provider: Option<String>,
114 pub model_id: Option<String>,
115 pub model_artifact_sha256: Option<String>,
116 pub dimensions: Option<usize>,
117 pub available: bool,
118 pub degraded: bool,
119 pub disabled: bool,
120 pub unavailable_reason: Option<String>,
121 pub provider_config: ProviderConfigSummary,
122 pub cold_start_embedding_latency_ms: Option<f64>,
123 pub query_embedding_latency_p95_ms: Option<f64>,
124 pub query_embedding_latency_samples: usize,
125 pub overall: Option<CategoryEvaluation>,
126 pub existing_slices: Option<CategoryEvaluation>,
127 pub existing_slice_details: BTreeMap<String, CategoryEvaluation>,
128 pub provider_comparison_slice: Option<CategoryEvaluation>,
129 pub query_summaries: Vec<ProviderQuerySummary>,
130}
131
132#[derive(Debug, Clone, Serialize)]
133pub struct ProviderConfigSummary {
134 pub provider: String,
135 pub fallback: Option<String>,
136 pub model: String,
137 pub base_url: String,
138 pub dimensions: Option<usize>,
139 pub api_key_env: String,
140 pub model_dir: Option<String>,
141 pub timeout_secs: u64,
142 pub api_calls_allowed: bool,
143}
144
145#[derive(Debug, Clone, Serialize)]
146pub struct ProviderQuerySummary {
147 pub id: String,
148 pub slice: String,
149 pub status: String,
150 pub result_count: usize,
151 pub retrieved_ids: Vec<i64>,
152 pub matched_refs: usize,
153 pub expected_refs: usize,
154 pub retrieval_latency_ms: f64,
155 pub query_embedding_latency_ms: Option<f64>,
156}
157
158#[derive(Debug, Clone, Serialize)]
159pub struct DefaultDecision {
160 pub change_default: bool,
161 pub decision: DefaultDecisionKind,
162 pub decision_reason: String,
163 pub criteria: DefaultFlipCriteria,
164 pub blockers: Vec<String>,
165}
166
167#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
168#[serde(rename_all = "snake_case")]
169pub enum DefaultDecisionKind {
170 KeepFeatureHash,
171 FlipToLocal,
172}
173
174#[derive(Debug, Clone, Serialize)]
175pub struct DefaultFlipCriteria {
176 pub local_available: bool,
177 pub api_reference_available: bool,
178 pub provider_comparison_slice_present: bool,
179 pub provider_comparison_slice_improves: bool,
180 pub paraphrase_slice_improves: bool,
181 pub existing_slices_within_budget: bool,
182 pub cold_start_embedding_latency_within_budget: bool,
183 pub query_embedding_latency_within_budget: bool,
184}
185
186pub fn run_provider_comparison_eval(
187 options: ProviderComparisonOptions,
188) -> Result<ProviderComparisonReport> {
189 let dataset = golden::load_dataset(&options.dataset_path)?;
190 run_provider_comparison_dataset(options, dataset)
191}
192
193fn run_provider_comparison_dataset(
194 options: ProviderComparisonOptions,
195 dataset: GoldenDataset,
196) -> Result<ProviderComparisonReport> {
197 let _env_guard = crate::runtime_config::ENV_LOCK
198 .lock()
199 .map_err(|_| anyhow!("embedding provider env lock poisoned"))?;
200 run_provider_comparison_dataset_locked(options, dataset)
201}
202
203fn run_provider_comparison_dataset_locked(
204 options: ProviderComparisonOptions,
205 dataset: GoldenDataset,
206) -> Result<ProviderComparisonReport> {
207 if !dataset.has_fixture_corpus() {
208 bail!("provider comparison requires a fixture-backed golden dataset");
209 }
210 ensure_provider_comparison_slice(&dataset)?;
211
212 let k = options.k.max(1);
213 let base_config = embedding::resolve_embedding_config()?;
214 let providers = vec![
215 evaluate_provider(
216 &dataset,
217 &options.dataset_path,
218 k,
219 &base_config,
220 EmbeddingProvider::FeatureHash,
221 options.allow_api,
222 )?,
223 evaluate_provider(
224 &dataset,
225 &options.dataset_path,
226 k,
227 &base_config,
228 EmbeddingProvider::Local,
229 options.allow_api,
230 )?,
231 evaluate_provider(
232 &dataset,
233 &options.dataset_path,
234 k,
235 &base_config,
236 EmbeddingProvider::OpenAi,
237 options.allow_api,
238 )?,
239 ];
240 let default_decision = build_default_decision(&providers);
241 Ok(ProviderComparisonReport {
242 version: REPORT_VERSION,
243 generated_at_epoch: chrono::Utc::now().timestamp(),
244 build_profile: if cfg!(debug_assertions) {
245 "debug"
246 } else {
247 "release"
248 },
249 target_os: std::env::consts::OS,
250 target_arch: std::env::consts::ARCH,
251 dataset_path: options.dataset_path,
252 k,
253 required_providers: vec!["feature-hash", "local", "api"],
254 provider_comparison_slice: PROVIDER_COMPARISON_SLICE,
255 query_embedding_latency_budget_p95_ms: QUERY_EMBEDDING_LATENCY_BUDGET_P95_MS,
256 cold_start_embedding_latency_budget_ms: COLD_START_EMBEDDING_LATENCY_BUDGET_MS,
257 existing_regression_budget: EXISTING_REGRESSION_BUDGET,
258 providers,
259 default_decision,
260 notes: vec![
261 "Provider rows are forced without fallback so unavailable providers cannot pass by silently using another embedding space.",
262 "Cold-start embedding latency measures provider verification plus the first profile probe before corpus seeding; query-embedding p95 measures subsequent warm-session queries.",
263 "Configured model-directory paths are redacted from committed evidence; model identity is recorded by model id and artifact SHA-256.",
264 "Remote API calls are opt-in with --allow-api; the default reference run records API as unavailable instead of spending network/API budget.",
265 "This report is evidence for GH-716 and does not change the default provider by itself.",
266 ],
267 })
268}
269
270fn ensure_provider_comparison_slice(dataset: &GoldenDataset) -> Result<()> {
271 let count = dataset
272 .queries
273 .iter()
274 .filter(|query| query.slice_label() == PROVIDER_COMPARISON_SLICE)
275 .count();
276 if count == 0 {
277 bail!("provider comparison requires at least one provider_comparison golden query");
278 }
279 Ok(())
280}
281
282fn evaluate_provider(
283 dataset: &GoldenDataset,
284 dataset_path: &str,
285 k: usize,
286 base_config: &EmbeddingConfig,
287 provider: EmbeddingProvider,
288 allow_api: bool,
289) -> Result<ProviderComparisonRow> {
290 let forced_config = forced_provider_config(base_config, provider)?;
291 if provider == EmbeddingProvider::OpenAi && !allow_api {
292 return Ok(unavailable_row(
293 provider,
294 &forced_config,
295 "api provider comparison skipped because --allow-api was not set",
296 allow_api,
297 ));
298 }
299
300 let _scope = ScopedEmbeddingConfig::activate(&forced_config, provider, allow_api)?;
301 let cold_start_started = Instant::now();
302 let status = embedding::embedding_provider_status_without_probe()
303 .with_context(|| format!("resolve {} provider status", provider.label()))?;
304 if let Some(reason) = status.unavailable_reason.clone() {
305 ensure_optional_provider(provider, reason.as_str())?;
306 return Ok(row_from_status_unavailable(
307 provider,
308 &forced_config,
309 status,
310 reason,
311 allow_api,
312 ));
313 }
314 if status.disabled {
315 ensure_optional_provider(provider, "embedding provider is disabled")?;
316 return Ok(row_from_status_unavailable(
317 provider,
318 &forced_config,
319 status,
320 "embedding provider is disabled".to_string(),
321 allow_api,
322 ));
323 }
324
325 let evaluated = model_pin::with_provider_model_pin(provider, &forced_config, || {
326 let active_profile = embedding::configured_backfill_target()
327 .with_context(|| format!("probe {} embedding profile", provider.label()))?;
328 let cold_start_embedding_latency_ms = cold_start_started.elapsed().as_secs_f64() * 1000.0;
329 let model_artifact_sha256 = match provider {
330 EmbeddingProvider::Local => Some(
331 embedding::configured_local_embedding_artifact_sha256(&forced_config)
332 .context("resolve pinned local embedding artifact digest")?,
333 ),
334 _ => None,
335 };
336 let evaluation = evaluate_available_provider(dataset, k)?;
337 Ok(row_from_evaluation(
338 provider,
339 &forced_config,
340 status.clone(),
341 active_profile.model,
342 active_profile.dimensions,
343 model_artifact_sha256,
344 cold_start_embedding_latency_ms,
345 evaluation,
346 allow_api,
347 ))
348 });
349 match evaluated {
350 Ok(row) => Ok(row),
351 Err(error) => {
352 let reason = format!("provider comparison failed for {dataset_path}: {error}");
353 optional_provider_error_row(provider, &forced_config, status, reason, allow_api)
354 }
355 }
356}
357
358fn ensure_optional_provider(provider: EmbeddingProvider, reason: &str) -> Result<()> {
359 if provider == EmbeddingProvider::FeatureHash {
360 bail!("feature-hash provider comparison baseline must be runnable: {reason}");
361 }
362 Ok(())
363}
364
365fn optional_provider_error_row(
366 provider: EmbeddingProvider,
367 config: &EmbeddingConfig,
368 status: EmbeddingProviderStatus,
369 reason: String,
370 allow_api: bool,
371) -> Result<ProviderComparisonRow> {
372 ensure_optional_provider(provider, &reason)?;
373 Ok(row_from_status_unavailable(
374 provider, config, status, reason, allow_api,
375 ))
376}
377
378pub(in crate::eval) fn forced_provider_config(
379 base_config: &EmbeddingConfig,
380 provider: EmbeddingProvider,
381) -> Result<EmbeddingConfig> {
382 let defaults = EmbeddingConfig::default();
383 let mut config = base_config.clone();
384 config.provider = provider;
385 config.fallback = None;
386 match provider {
387 EmbeddingProvider::FeatureHash => {
388 config.model = embedding::FEATURE_HASH_EMBEDDING_MODEL.to_string();
389 config.dimensions = Some(embedding::FEATURE_HASH_EMBEDDING_DIMENSIONS);
390 }
391 EmbeddingProvider::Local => {
392 if base_config.provider != EmbeddingProvider::Local {
393 config.model = defaults.model.clone();
394 }
395 config.model = embedding::configured_local_embedding_model_id(&config)
396 .context("resolve forced local provider model id")?;
397 config.dimensions = None;
398 }
399 EmbeddingProvider::OpenAi => {
400 if !matches!(
401 base_config.provider,
402 EmbeddingProvider::Auto | EmbeddingProvider::OpenAi
403 ) {
404 config.model = defaults.model.clone();
405 config.base_url = defaults.base_url.clone();
406 config.dimensions = defaults.dimensions;
407 }
408 }
409 EmbeddingProvider::Auto | EmbeddingProvider::Off => {}
410 }
411 Ok(config)
412}
413
414fn unavailable_row(
415 provider: EmbeddingProvider,
416 config: &EmbeddingConfig,
417 reason: impl Into<String>,
418 allow_api: bool,
419) -> ProviderComparisonRow {
420 let reason = reason.into();
421 let reason = redact_provider_reason(config, &reason);
422 ProviderComparisonRow {
423 provider: provider.label(),
424 configured_provider: provider.label().to_string(),
425 active_provider: provider.label().to_string(),
426 fallback_provider: None,
427 model_id: configured_model_id(provider, config),
428 model_artifact_sha256: None,
429 dimensions: config.dimensions,
430 available: false,
431 degraded: false,
432 disabled: false,
433 unavailable_reason: Some(reason),
434 provider_config: ProviderConfigSummary::from_config(config, allow_api),
435 cold_start_embedding_latency_ms: None,
436 query_embedding_latency_p95_ms: None,
437 query_embedding_latency_samples: 0,
438 overall: None,
439 existing_slices: None,
440 existing_slice_details: BTreeMap::new(),
441 provider_comparison_slice: None,
442 query_summaries: vec![],
443 }
444}
445
446fn row_from_status_unavailable(
447 provider: EmbeddingProvider,
448 config: &EmbeddingConfig,
449 status: EmbeddingProviderStatus,
450 reason: String,
451 allow_api: bool,
452) -> ProviderComparisonRow {
453 let reason = redact_provider_reason(config, &reason);
454 ProviderComparisonRow {
455 provider: provider.label(),
456 configured_provider: status.configured_provider,
457 active_provider: status.active_provider,
458 fallback_provider: status.fallback_provider,
459 model_id: status
460 .active_model_id
461 .or_else(|| configured_model_id(provider, config)),
462 model_artifact_sha256: None,
463 dimensions: status.active_dimensions.or(config.dimensions),
464 available: false,
465 degraded: status.degraded,
466 disabled: status.disabled,
467 unavailable_reason: Some(reason),
468 provider_config: ProviderConfigSummary::from_config(config, allow_api),
469 cold_start_embedding_latency_ms: None,
470 query_embedding_latency_p95_ms: None,
471 query_embedding_latency_samples: 0,
472 overall: None,
473 existing_slices: None,
474 existing_slice_details: BTreeMap::new(),
475 provider_comparison_slice: None,
476 query_summaries: vec![],
477 }
478}
479
480fn row_from_evaluation(
481 provider: EmbeddingProvider,
482 config: &EmbeddingConfig,
483 status: EmbeddingProviderStatus,
484 active_model_id: String,
485 active_dimensions: usize,
486 model_artifact_sha256: Option<String>,
487 cold_start_embedding_latency_ms: f64,
488 evaluation: ProviderRunEvaluation,
489 allow_api: bool,
490) -> ProviderComparisonRow {
491 let query_embedding_latency_p95_ms = (!evaluation.query_embedding_latencies_ms.is_empty())
492 .then(|| golden::run::percentile(evaluation.query_embedding_latencies_ms.clone(), 95.0));
493 ProviderComparisonRow {
494 provider: provider.label(),
495 configured_provider: status.configured_provider,
496 active_provider: status.active_provider,
497 fallback_provider: status.fallback_provider,
498 model_id: Some(active_model_id),
499 model_artifact_sha256,
500 dimensions: Some(active_dimensions),
501 available: true,
502 degraded: status.degraded,
503 disabled: status.disabled,
504 unavailable_reason: None,
505 provider_config: ProviderConfigSummary::from_config(config, allow_api),
506 cold_start_embedding_latency_ms: Some(cold_start_embedding_latency_ms),
507 query_embedding_latency_p95_ms,
508 query_embedding_latency_samples: evaluation.query_embedding_latencies_ms.len(),
509 overall: Some(evaluation.overall),
510 existing_slices: Some(evaluation.existing_slices),
511 existing_slice_details: evaluation.existing_slice_details,
512 provider_comparison_slice: Some(evaluation.provider_comparison_slice),
513 query_summaries: evaluation.query_summaries,
514 }
515}
516
517fn redact_provider_reason(config: &EmbeddingConfig, reason: &str) -> String {
518 let mut redacted = reason.to_string();
519 let mut sensitive_paths = Vec::new();
520 if let Ok(model_root) = embedding::configured_local_embedding_model_root(config) {
521 sensitive_paths.push(model_root.clone());
522 if let Ok(canonical_root) = std::fs::canonicalize(&model_root) {
523 sensitive_paths.push(canonical_root);
524 }
525 }
526 if let Some(configured) = config.model_dir.as_deref() {
527 sensitive_paths.push(PathBuf::from(configured));
528 }
529 sensitive_paths.sort();
530 sensitive_paths.dedup();
531 for path in sensitive_paths {
532 let displayed = path.to_string_lossy();
533 if !displayed.is_empty() {
534 redacted = redacted.replace(displayed.as_ref(), "<model-dir>");
535 }
536 }
537 redacted
538}
539
540fn configured_model_id(provider: EmbeddingProvider, config: &EmbeddingConfig) -> Option<String> {
541 match provider {
542 EmbeddingProvider::FeatureHash => Some(embedding::FEATURE_HASH_EMBEDDING_MODEL.to_string()),
543 EmbeddingProvider::OpenAi => Some(config.model.clone()),
544 EmbeddingProvider::Local => embedding::configured_local_embedding_model_id(config).ok(),
545 EmbeddingProvider::Auto | EmbeddingProvider::Off => None,
546 }
547}
548
549impl ProviderConfigSummary {
550 fn from_config(config: &EmbeddingConfig, allow_api: bool) -> Self {
551 Self {
552 provider: config.provider.label().to_string(),
553 fallback: config.fallback.map(|provider| provider.label().to_string()),
554 model: config.model.clone(),
555 base_url: config.base_url.clone(),
556 dimensions: config.dimensions,
557 api_key_env: config.api_key_env.clone(),
558 model_dir: config
559 .model_dir
560 .as_ref()
561 .map(|_| "<configured>".to_string()),
562 timeout_secs: config.timeout_secs,
563 api_calls_allowed: allow_api,
564 }
565 }
566}
567
568struct ProviderRunEvaluation {
569 overall: CategoryEvaluation,
570 existing_slices: CategoryEvaluation,
571 existing_slice_details: BTreeMap<String, CategoryEvaluation>,
572 provider_comparison_slice: CategoryEvaluation,
573 query_embedding_latencies_ms: Vec<f64>,
574 query_summaries: Vec<ProviderQuerySummary>,
575}
576
577fn evaluate_available_provider(dataset: &GoldenDataset, k: usize) -> Result<ProviderRunEvaluation> {
578 let conn = Connection::open_in_memory().context("open in-memory provider comparison DB")?;
579 crate::migrate::run_migrations(&conn).context("migrate provider comparison DB")?;
580 golden::run::seed_fixture_corpus(&conn, &dataset.corpus)
581 .context("seed provider comparison fixture corpus")?;
582
583 let mut overall = golden::run::CategoryAccumulator::default();
584 let mut existing_slices = golden::run::CategoryAccumulator::default();
585 let mut existing_slice_details = BTreeMap::<String, golden::run::CategoryAccumulator>::new();
586 let mut provider_comparison_slice = golden::run::CategoryAccumulator::default();
587 let mut query_embedding_latencies_ms = Vec::new();
588 let mut query_summaries = Vec::with_capacity(dataset.queries.len());
589
590 for query in &dataset.queries {
591 let started = Instant::now();
592 let (results, explain) = crate::retrieval::search::search_with_branch_explain(
593 &conn,
594 Some(&query.query),
595 query.project.as_deref(),
596 query.memory_type.as_deref(),
597 k.max(10) as i64,
598 0,
599 false,
600 query.branch.as_deref(),
601 )?;
602 let retrieval_latency_ms = started.elapsed().as_secs_f64() * 1000.0;
603 let query_embedding_latency_ms = explain
604 .as_ref()
605 .and_then(|explain| phase_latency_ms(&explain.timings, "query_embedding"));
606 if let Some(latency_ms) = query_embedding_latency_ms {
607 query_embedding_latencies_ms.push(latency_ms);
608 }
609 let query_tokens = golden::run::estimate_query_tokens(&query.query);
610 let evaluation =
611 golden::run::evaluate_query(query, &results, k, query_tokens, retrieval_latency_ms);
612
613 golden::run::record_bucket(&mut overall, query, &evaluation);
614 if query.slice_label() == PROVIDER_COMPARISON_SLICE {
615 golden::run::record_bucket(&mut provider_comparison_slice, query, &evaluation);
616 } else {
617 golden::run::record_bucket(&mut existing_slices, query, &evaluation);
618 golden::run::record_bucket(
619 existing_slice_details
620 .entry(query.slice_label().to_string())
621 .or_default(),
622 query,
623 &evaluation,
624 );
625 }
626
627 query_summaries.push(ProviderQuerySummary {
628 id: evaluation.id,
629 slice: evaluation.slice,
630 status: evaluation.status.label().to_string(),
631 result_count: evaluation.result_count,
632 retrieved_ids: evaluation.retrieved_ids,
633 matched_refs: evaluation.matched_refs,
634 expected_refs: evaluation.expected_refs,
635 retrieval_latency_ms,
636 query_embedding_latency_ms,
637 });
638 }
639
640 Ok(ProviderRunEvaluation {
641 overall: golden::run::bucket_evaluation(overall),
642 existing_slices: golden::run::bucket_evaluation(existing_slices),
643 existing_slice_details: existing_slice_details
644 .into_iter()
645 .map(|(slice, bucket)| (slice, golden::run::bucket_evaluation(bucket)))
646 .collect(),
647 provider_comparison_slice: golden::run::bucket_evaluation(provider_comparison_slice),
648 query_embedding_latencies_ms,
649 query_summaries,
650 })
651}
652
653fn phase_latency_ms(timings: &[crate::perf::PhaseTiming], phase: &str) -> Option<f64> {
654 timings
655 .iter()
656 .find(|timing| timing.phase == phase)
657 .map(|timing| timing.elapsed_ms as f64)
658}
659
660pub(in crate::eval) struct ScopedEmbeddingConfig {
661 saved: Vec<(String, Option<String>)>,
662 config_path: PathBuf,
663}
664
665impl ScopedEmbeddingConfig {
666 pub(in crate::eval) fn activate(
667 config: &EmbeddingConfig,
668 provider: EmbeddingProvider,
669 allow_api: bool,
670 ) -> Result<Self> {
671 let mut keys = BASE_ENV_KEYS
672 .iter()
673 .map(|key| (*key).to_string())
674 .collect::<Vec<_>>();
675 if !keys.iter().any(|key| key == &config.api_key_env) {
676 keys.push(config.api_key_env.clone());
677 }
678 let saved = keys
679 .iter()
680 .map(|key| (key.clone(), std::env::var(key).ok()))
681 .collect::<Vec<_>>();
682 let config_path = temp_config_path(provider);
683 write_temp_embedding_config(&config_path, config, provider)?;
684
685 for key in &keys {
686 unsafe { std::env::remove_var(key) };
687 }
688 unsafe {
689 std::env::set_var(ENV_CONFIG, &config_path);
690 }
691 if provider == EmbeddingProvider::OpenAi && allow_api {
692 restore_api_key_for_scoped_config(config, &saved);
693 }
694
695 Ok(Self { saved, config_path })
696 }
697}
698
699impl Drop for ScopedEmbeddingConfig {
700 fn drop(&mut self) {
701 for (key, value) in self.saved.drain(..) {
702 match value {
703 Some(value) => unsafe { std::env::set_var(key, value) },
704 None => unsafe { std::env::remove_var(key) },
705 }
706 }
707 let _ = fs::remove_file(&self.config_path);
708 }
709}
710
711fn restore_api_key_for_scoped_config(config: &EmbeddingConfig, saved: &[(String, Option<String>)]) {
712 let saved_values = saved
713 .iter()
714 .filter_map(|(key, value)| value.as_ref().map(|value| (key.as_str(), value.as_str())))
715 .collect::<BTreeMap<_, _>>();
716 if let Some(value) = saved_values.get(ENV_API_KEY) {
717 unsafe { std::env::set_var(ENV_API_KEY, value) };
718 } else if let Some(value) = saved_values.get(ENV_API_KEY_LEGACY) {
719 unsafe { std::env::set_var(ENV_API_KEY_LEGACY, value) };
720 } else if let Some(value) = saved_values.get(config.api_key_env.as_str()) {
721 unsafe { std::env::set_var(&config.api_key_env, value) };
722 } else if let Some(value) = saved_values.get(DEFAULT_API_KEY_ENV) {
723 unsafe { std::env::set_var(DEFAULT_API_KEY_ENV, value) };
724 }
725}
726
727fn temp_config_path(provider: EmbeddingProvider) -> PathBuf {
728 std::env::temp_dir().join(format!(
729 "remem-provider-comparison-{}-{}-{}.toml",
730 std::process::id(),
731 chrono::Utc::now().timestamp_nanos_opt().unwrap_or_default(),
732 provider.label()
733 ))
734}
735
736fn write_temp_embedding_config(
737 path: &Path,
738 config: &EmbeddingConfig,
739 provider: EmbeddingProvider,
740) -> Result<()> {
741 let mut doc = DocumentMut::new();
742 let mut table = Table::new();
743 table["provider"] = value(provider.label());
744 table["model"] = value(config.model.clone());
745 table["base_url"] = value(config.base_url.clone());
746 if let Some(dimensions) = config.dimensions {
747 table["dimensions"] = value(dimensions as i64);
748 }
749 table["api_key_env"] = value(config.api_key_env.clone());
750 if let Some(model_dir) = config.model_dir.as_deref() {
751 table["model_dir"] = value(model_dir);
752 }
753 table["timeout_secs"] = value(config.timeout_secs as i64);
754 doc["embeddings"] = Item::Table(table);
755 fs::write(path, doc.to_string())
756 .with_context(|| format!("write provider comparison config {}", path.display()))
757}
758
759#[cfg(test)]
760mod tests;