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