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;