Skip to main content

remem/eval/
provider_comparison.rs

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;