1use zeph_llm::any::AnyProvider;
12use zeph_llm::claude::ClaudeProvider;
13#[cfg(feature = "cocoon")]
14use zeph_llm::cocoon::{CocoonClient, CocoonProvider};
15use zeph_llm::compatible::CompatibleProvider;
16use zeph_llm::gemini::GeminiProvider;
17#[cfg(feature = "gonka")]
18use zeph_llm::gonka::endpoints::{EndpointPool, GonkaEndpoint};
19#[cfg(feature = "gonka")]
20use zeph_llm::gonka::{GonkaProvider, RequestSigner};
21use zeph_llm::http::llm_client;
22use zeph_llm::ollama::OllamaProvider;
23use zeph_llm::openai::OpenAiProvider;
24#[cfg(feature = "gonka")]
25use zeroize::Zeroizing;
26
27use crate::agent::state::ProviderConfigSnapshot;
28#[cfg(feature = "candle")]
29use crate::config::{CandleDevice, CandleInlineConfig, CandleSource};
30use crate::config::{Config, ProviderEntry, ProviderKind};
31
32#[non_exhaustive]
33#[derive(Debug, thiserror::Error)]
40pub enum BootstrapError {
41 #[error("config error: {0}")]
43 Config(#[from] crate::config::ConfigError),
44 #[error("provider error: {0}")]
46 Provider(String),
47 #[error("memory error: {0}")]
49 Memory(String),
50 #[error("vault init error: {0}")]
52 VaultInit(crate::vault::AgeVaultError),
53 #[error("I/O error: {0}")]
55 Io(#[from] std::io::Error),
56}
57
58pub fn build_provider_for_switch(
75 entry: &ProviderEntry,
76 snapshot: &ProviderConfigSnapshot,
77 secret_registry: Option<&std::sync::Arc<zeph_sanitizer::secret_mask::SecretMaskRegistry>>,
78) -> Result<AnyProvider, BootstrapError> {
79 use zeph_common::secret::Secret;
80 let mut config = Config::default();
84 config.secrets.claude_api_key = snapshot.claude_api_key.as_deref().map(Secret::new);
85 config.secrets.openai_api_key = snapshot.openai_api_key.as_deref().map(Secret::new);
86 config.secrets.gemini_api_key = snapshot.gemini_api_key.as_deref().map(Secret::new);
87 config.secrets.compatible_api_keys = snapshot
88 .compatible_api_keys
89 .iter()
90 .map(|(k, v)| (k.clone(), Secret::new(v.as_str())))
91 .collect();
92 config.secrets.gonka_private_key = snapshot
93 .gonka_private_key
94 .as_ref()
95 .map(|z| Secret::new(z.as_str()));
96 config.secrets.gonka_address = snapshot.gonka_address.as_deref().map(Secret::new);
97 config.secrets.cocoon_access_hash = snapshot.cocoon_access_hash.as_deref().map(Secret::new);
98 config.timeouts.llm_request_timeout_secs = snapshot.llm_request_timeout_secs;
99 config
100 .llm
101 .embedding_model
102 .clone_from(&snapshot.embedding_model);
103 build_provider_from_entry(entry, &config, secret_registry)
104}
105
106pub fn build_provider_from_entry(
125 entry: &ProviderEntry,
126 config: &Config,
127 secret_registry: Option<&std::sync::Arc<zeph_sanitizer::secret_mask::SecretMaskRegistry>>,
128) -> Result<AnyProvider, BootstrapError> {
129 let provider = build_provider_from_entry_inner(entry, config)?;
130 Ok(match secret_registry {
131 Some(registry) => provider.masked(std::sync::Arc::clone(registry)
132 as std::sync::Arc<dyn zeph_llm::masking::OutboundMasker>),
133 None => provider,
134 })
135}
136
137fn build_provider_from_entry_inner(
138 entry: &ProviderEntry,
139 config: &Config,
140) -> Result<AnyProvider, BootstrapError> {
141 match entry.provider_type {
142 ProviderKind::Ollama => Ok(build_ollama_provider(entry, config)),
143 ProviderKind::Claude => build_claude_provider(entry, config),
144 ProviderKind::OpenAi => build_openai_provider(entry, config),
145 ProviderKind::Gemini => build_gemini_provider(entry, config),
146 ProviderKind::Compatible => build_compatible_provider(entry, config),
147 #[cfg(feature = "candle")]
148 ProviderKind::Candle => build_candle_provider(entry, config),
149 #[cfg(not(feature = "candle"))]
150 ProviderKind::Candle => Err(BootstrapError::Provider(
151 "candle feature is not enabled".into(),
152 )),
153 #[cfg(feature = "gonka")]
154 ProviderKind::Gonka => build_gonka_provider(entry, config),
155 #[cfg(not(feature = "gonka"))]
156 ProviderKind::Gonka => Err(BootstrapError::Provider(
157 "gonka feature is not enabled; rebuild with --features gonka".into(),
158 )),
159 #[cfg(feature = "cocoon")]
160 ProviderKind::Cocoon => build_cocoon_provider(entry, config),
161 #[cfg(not(feature = "cocoon"))]
162 ProviderKind::Cocoon => Err(BootstrapError::Provider(
163 "cocoon feature is not enabled; rebuild with --features cocoon".into(),
164 )),
165 _ => Err(BootstrapError::Provider(format!(
166 "unknown provider kind: {:?}",
167 entry.provider_type
168 ))),
169 }
170}
171
172#[must_use]
187pub fn resolve_named_provider(config: &Config, primary: &AnyProvider, name: &str) -> AnyProvider {
188 if name.is_empty() {
189 return primary.clone();
190 }
191 let Some(entry) = config
192 .llm
193 .providers
194 .iter()
195 .find(|e| e.effective_name().eq_ignore_ascii_case(name))
196 else {
197 tracing::warn!(
198 provider = name,
199 "provider not found in [[llm.providers]], falling back to primary"
200 );
201 return primary.clone();
202 };
203 match build_provider_from_entry(entry, config, None) {
204 Ok(provider) => provider,
205 Err(e) => {
206 tracing::warn!(error = %e, provider = name, "failed to build named provider, falling back to primary");
207 primary.clone()
208 }
209 }
210}
211
212#[must_use]
225pub fn build_resume_condenser(
226 config: &Config,
227 primary_provider: &AnyProvider,
228) -> (
229 zeph_session::LlmCondenser,
230 std::sync::Arc<zeph_agent_context::memory_backend::TokenCounterAdapter>,
231) {
232 let condense_config = &config.session.condense;
233 let condense_provider = resolve_named_provider(
234 config,
235 primary_provider,
236 condense_config.condense_provider.as_str(),
237 );
238 let token_counter_adapter = std::sync::Arc::new(
239 zeph_agent_context::memory_backend::TokenCounterAdapter::new(std::sync::Arc::new(
240 zeph_memory::TokenCounter::new(),
241 )),
242 );
243 let condenser = zeph_session::LlmCondenser::new(
244 zeph_context::summarization::SummarizationDeps {
245 provider: condense_provider,
246 llm_timeout: std::time::Duration::from_secs(config.timeouts.llm_seconds),
247 token_counter: std::sync::Arc::new(
248 zeph_agent_context::memory_backend::TokenCounterAdapter::new(std::sync::Arc::new(
249 zeph_memory::TokenCounter::new(),
250 )),
251 ),
252 structured_summaries: config.memory.structured_summaries,
253 on_anchored_summary: None,
254 },
255 condense_config.threshold,
256 condense_config.keep_recent,
257 );
258 (condenser, token_counter_adapter)
259}
260
261fn build_ollama_provider(entry: &ProviderEntry, config: &Config) -> AnyProvider {
262 let base_url = entry
263 .base_url
264 .as_deref()
265 .unwrap_or("http://localhost:11434");
266 let model = entry.model.as_deref().unwrap_or("qwen3:8b").to_owned();
267 let embed = entry
268 .embedding_model
269 .clone()
270 .unwrap_or_else(|| config.llm.embedding_model.clone());
271 let mut provider = OllamaProvider::new(base_url, model, embed);
272 if let Some(ref vm) = entry.vision_model {
273 provider = provider.with_vision_model(vm.clone());
274 }
275 if config.mcp.forward_output_schema {
276 tracing::debug!(
277 "mcp.forward_output_schema is enabled but Ollama does not support \
278 output schema forwarding; setting ignored for this provider"
279 );
280 }
281 AnyProvider::Ollama(provider)
282}
283
284fn build_claude_provider(
285 entry: &ProviderEntry,
286 config: &Config,
287) -> Result<AnyProvider, BootstrapError> {
288 let api_key = config
289 .secrets
290 .claude_api_key
291 .as_ref()
292 .ok_or_else(|| BootstrapError::Provider("ZEPH_CLAUDE_API_KEY not found in vault".into()))?
293 .expose()
294 .to_owned();
295 let model = entry
296 .model
297 .clone()
298 .unwrap_or_else(|| "claude-haiku-4-5-20251001".to_owned());
299 let max_tokens = entry.max_tokens.unwrap_or(4096);
300 let provider = ClaudeProvider::new(api_key, model, max_tokens)
301 .with_client(llm_client(config.timeouts.llm_request_timeout_secs))
302 .with_extended_context(entry.enable_extended_context)
303 .with_thinking_opt(entry.thinking.clone())
304 .map_err(|e| BootstrapError::Provider(format!("invalid thinking config: {e}")))?
305 .with_server_compaction(entry.server_compaction)
306 .with_prompt_cache_ttl(entry.prompt_cache_ttl)
307 .with_stream_limits(config.llm.stream_limits.clone())
308 .with_output_schema_forwarding(
309 config.mcp.forward_output_schema,
310 config.mcp.output_schema_hint_bytes,
311 config.mcp.max_description_bytes,
312 );
313 tracing::info!(
314 forward = config.mcp.forward_output_schema,
315 "mcp.output_schema.forwarding_configured"
316 );
317 Ok(AnyProvider::Claude(provider))
318}
319
320fn build_openai_provider(
321 entry: &ProviderEntry,
322 config: &Config,
323) -> Result<AnyProvider, BootstrapError> {
324 let api_key = config
325 .secrets
326 .openai_api_key
327 .as_ref()
328 .ok_or_else(|| BootstrapError::Provider("ZEPH_OPENAI_API_KEY not found in vault".into()))?
329 .expose()
330 .to_owned();
331 let base_url = entry
332 .base_url
333 .clone()
334 .unwrap_or_else(|| "https://api.openai.com/v1".to_owned());
335 let model = entry
336 .model
337 .clone()
338 .unwrap_or_else(|| "gpt-4o-mini".to_owned());
339 let max_tokens = entry.max_tokens.unwrap_or(4096);
340 Ok(AnyProvider::OpenAi(
341 OpenAiProvider::new(zeph_llm::OpenAiConfig {
342 api_key,
343 base_url,
344 model,
345 max_tokens,
346 embedding_model: entry.embedding_model.clone(),
347 reasoning_effort: entry.reasoning_effort.clone(),
348 context_window: None,
349 completion_tokens_param: None,
350 })
351 .with_client(llm_client(config.timeouts.llm_request_timeout_secs))
352 .with_output_schema_forwarding(
353 config.mcp.forward_output_schema,
354 config.mcp.output_schema_hint_bytes,
355 config.mcp.max_description_bytes,
356 ),
357 ))
358}
359
360fn build_gemini_provider(
361 entry: &ProviderEntry,
362 config: &Config,
363) -> Result<AnyProvider, BootstrapError> {
364 let api_key = config
365 .secrets
366 .gemini_api_key
367 .as_ref()
368 .ok_or_else(|| BootstrapError::Provider("ZEPH_GEMINI_API_KEY not found in vault".into()))?
369 .expose()
370 .to_owned();
371 let model = entry
372 .model
373 .clone()
374 .unwrap_or_else(|| "gemini-2.0-flash".to_owned());
375 let max_tokens = entry.max_tokens.unwrap_or(8192);
376 let base_url = entry
377 .base_url
378 .clone()
379 .unwrap_or_else(|| "https://generativelanguage.googleapis.com".to_owned());
380 let mut provider = GeminiProvider::new(api_key, model, max_tokens)
381 .with_base_url(base_url)
382 .with_client(llm_client(config.timeouts.llm_request_timeout_secs));
383 if let Some(ref em) = entry.embedding_model {
384 provider = provider.with_embedding_model(em.clone());
385 }
386 if let Some(level) = entry.thinking_level {
387 provider = provider.with_thinking_level(level);
388 }
389 if let Some(budget) = entry.thinking_budget {
390 provider = provider
391 .with_thinking_budget(budget)
392 .map_err(|e| BootstrapError::Provider(e.to_string()))?;
393 }
394 if let Some(include) = entry.include_thoughts {
395 provider = provider.with_include_thoughts(include);
396 }
397 if config.mcp.forward_output_schema {
398 tracing::debug!(
399 "mcp.forward_output_schema is enabled but Gemini does not support \
400 output schema forwarding; setting ignored for this provider"
401 );
402 }
403 Ok(AnyProvider::Gemini(provider))
404}
405
406fn build_compatible_provider(
407 entry: &ProviderEntry,
408 config: &Config,
409) -> Result<AnyProvider, BootstrapError> {
410 let name = entry.name.as_deref().ok_or_else(|| {
411 BootstrapError::Provider(
412 "compatible provider requires 'name' field in [[llm.providers]]".into(),
413 )
414 })?;
415 let base_url = entry.base_url.clone().ok_or_else(|| {
416 BootstrapError::Provider(format!("compatible provider '{name}' requires 'base_url'"))
417 })?;
418 let model = entry.model.clone().unwrap_or_default();
419 let api_key = entry.api_key.clone().unwrap_or_else(|| {
420 config
421 .secrets
422 .compatible_api_keys
423 .get(name)
424 .map(|s| s.expose().to_owned())
425 .unwrap_or_default()
426 });
427 let max_tokens = entry.max_tokens.unwrap_or(4096);
428 let provider = CompatibleProvider::new(zeph_llm::CompatibleConfig {
429 provider_name: name.to_owned(),
430 api_key,
431 base_url,
432 model,
433 max_tokens,
434 embedding_model: entry.embedding_model.clone(),
435 completion_tokens_param: None,
436 })
437 .with_output_schema_forwarding(
438 config.mcp.forward_output_schema,
439 config.mcp.output_schema_hint_bytes,
440 config.mcp.max_description_bytes,
441 );
442 tracing::info!(
443 forward = config.mcp.forward_output_schema,
444 provider = name,
445 "mcp.output_schema.forwarding_configured"
446 );
447 Ok(AnyProvider::Compatible(provider))
448}
449
450#[cfg(feature = "gonka")]
451fn build_gonka_provider(
452 entry: &ProviderEntry,
453 config: &Config,
454) -> Result<AnyProvider, BootstrapError> {
455 let _span = tracing::info_span!("core.provider_factory.build_gonka").entered();
456
457 let private_key_hex: Zeroizing<String> = Zeroizing::new(
458 config
459 .secrets
460 .gonka_private_key
461 .as_ref()
462 .ok_or_else(|| {
463 BootstrapError::Provider(
464 "ZEPH_GONKA_PRIVATE_KEY not found in vault; set it with: zeph vault set ZEPH_GONKA_PRIVATE_KEY <hex>".into(),
465 )
466 })?
467 .expose()
468 .to_owned(),
469 );
470
471 let chain_prefix = entry.effective_gonka_chain_prefix().to_owned();
472 let signer = RequestSigner::from_hex(&private_key_hex, &chain_prefix)
473 .map_err(|e| BootstrapError::Provider(format!("invalid Gonka private key: {e}")))?;
474
475 if let Some(ref configured_address) = config.secrets.gonka_address {
476 let configured = configured_address.expose().to_lowercase();
477 let derived = signer.address().to_lowercase();
478 if configured != derived {
479 return Err(BootstrapError::Provider(format!(
480 "ZEPH_GONKA_ADDRESS does not match address derived from private key \
481 (configured: {configured}, derived: {derived})"
482 )));
483 }
484 } else {
485 tracing::info!(
486 address = signer.address(),
487 "Gonka: using address derived from private key (ZEPH_GONKA_ADDRESS not set)"
488 );
489 }
490
491 if entry.gonka_nodes.is_empty() {
492 return Err(BootstrapError::Provider(
493 "Gonka provider entry must have at least one node in gonka_nodes".into(),
494 ));
495 }
496
497 let endpoints: Vec<GonkaEndpoint> = entry
498 .gonka_nodes
499 .iter()
500 .map(|n| GonkaEndpoint {
501 base_url: n.url.clone(),
502 address: n.address.clone(),
503 })
504 .collect();
505
506 let pool = EndpointPool::new(endpoints).map_err(|e| {
507 BootstrapError::Provider(format!("failed to build Gonka endpoint pool: {e}"))
508 })?;
509
510 let model = entry.model.clone().unwrap_or_else(|| "gpt-4o".to_owned());
511 let max_tokens = entry.max_tokens.unwrap_or(4096);
512 let timeout = std::time::Duration::from_secs(config.timeouts.llm_request_timeout_secs);
513
514 let provider = GonkaProvider::new(zeph_llm::gonka::GonkaConfig {
515 signer: std::sync::Arc::new(signer),
516 pool: std::sync::Arc::new(pool),
517 model,
518 max_tokens,
519 embedding_model: entry.embedding_model.clone(),
520 timeout,
521 });
522
523 Ok(AnyProvider::Gonka(provider))
524}
525
526#[cfg(feature = "cocoon")]
530struct CocoonClientParams {
531 base_url: String,
532 access_hash: Option<String>,
533 timeout: std::time::Duration,
534}
535
536#[cfg(feature = "cocoon")]
548fn resolve_cocoon_client_params(
549 entry: &ProviderEntry,
550 config: &Config,
551) -> Result<CocoonClientParams, BootstrapError> {
552 let base_url = entry
553 .cocoon_client_url
554 .as_deref()
555 .unwrap_or("http://localhost:10000");
556
557 if !base_url.starts_with("http://localhost")
559 && !base_url.starts_with("http://127.0.0.1")
560 && !base_url.starts_with("http://[::1]")
561 && !base_url.starts_with("https://localhost")
562 && !base_url.starts_with("https://127.0.0.1")
563 && !base_url.starts_with("https://[::1]")
564 {
565 tracing::warn!(
566 url = base_url,
567 "cocoon_client_url points to a non-localhost host; \
568 ensure this is intentional (expected sidecar on localhost)"
569 );
570 }
571
572 if entry
573 .cocoon_access_hash
574 .as_deref()
575 .is_some_and(|v| !v.is_empty())
576 {
577 tracing::warn!(
578 "cocoon_access_hash in config file appears to contain a raw value; \
579 this field should be empty — the actual hash must be stored in the vault: \
580 zeph vault set ZEPH_COCOON_ACCESS_HASH <hash>"
581 );
582 }
583
584 let access_hash = if entry.cocoon_access_hash.is_some() {
585 let hash = config
586 .secrets
587 .cocoon_access_hash
588 .as_ref()
589 .ok_or_else(|| {
590 BootstrapError::Provider(
591 "ZEPH_COCOON_ACCESS_HASH not found in vault; set it with: \
592 zeph vault set ZEPH_COCOON_ACCESS_HASH <hash>"
593 .into(),
594 )
595 })?
596 .expose()
597 .to_owned();
598 Some(hash)
599 } else {
600 None
601 };
602
603 let timeout = std::time::Duration::from_secs(config.timeouts.llm_request_timeout_secs);
604
605 Ok(CocoonClientParams {
606 base_url: base_url.to_owned(),
607 access_hash,
608 timeout,
609 })
610}
611
612#[cfg(feature = "cocoon")]
619fn build_cocoon_provider(
620 entry: &ProviderEntry,
621 config: &Config,
622) -> Result<AnyProvider, BootstrapError> {
623 let _span = tracing::info_span!("core.provider_factory.build_cocoon").entered();
624
625 let params = resolve_cocoon_client_params(entry, config)?;
626 let client = std::sync::Arc::new(CocoonClient::new(
627 ¶ms.base_url,
628 params.access_hash,
629 params.timeout,
630 ));
631
632 let model = entry
633 .model
634 .clone()
635 .unwrap_or_else(|| "Qwen/Qwen3-0.6B".to_owned());
636 let max_tokens = entry.max_tokens.unwrap_or(4096);
637 let provider = CocoonProvider::new(model, max_tokens, entry.embedding_model.clone(), client);
638
639 Ok(AnyProvider::Cocoon(provider))
640}
641
642#[cfg(feature = "cocoon")]
656pub fn spawn_cocoon_health_checks(
657 providers: &[&ProviderEntry],
658 config: &Config,
659 supervisor: &std::sync::Arc<zeph_common::TaskSupervisor>,
660) {
661 for entry in providers {
662 if entry.provider_type != ProviderKind::Cocoon || !entry.cocoon_health_check {
663 continue;
664 }
665 let params = match resolve_cocoon_client_params(entry, config) {
666 Ok(params) => params,
667 Err(e) => {
668 tracing::warn!(
669 name = entry.name.as_deref().unwrap_or("<unnamed>"),
670 error = %e,
671 "skipping cocoon health check: failed to resolve client params"
672 );
673 continue;
674 }
675 };
676 let client = std::sync::Arc::new(CocoonClient::new(
677 ¶ms.base_url,
678 params.access_hash,
679 params.timeout,
680 ));
681 supervisor.spawn(zeph_common::task_supervisor::TaskDescriptor {
682 name: "core.provider_factory.cocoon_health_check",
683 restart: zeph_common::task_supervisor::RestartPolicy::RunOnce,
684 factory: move || {
685 let client = client.clone();
686 async move {
687 match client.health_check().await {
688 Ok(h) => {
689 tracing::info!(
690 proxy_connected = h.proxy_connected,
691 worker_count = h.worker_count,
692 "cocoon sidecar health check passed"
693 );
694 }
695 Err(e) => {
696 tracing::warn!(
697 error = %e,
698 "cocoon sidecar health check failed; \
699 inference requests will return LlmError::Unavailable until the sidecar is running"
700 );
701 }
702 }
703 }
704 },
705 });
706 }
707}
708
709#[cfg(feature = "candle")]
715struct CandleLoadParams {
716 source: zeph_llm::candle_provider::loader::ModelSource,
717 template: zeph_llm::candle_provider::template::ChatTemplate,
718 gen_config: zeph_llm::candle_provider::generate::GenerationConfig,
719 embedding_repo: Option<String>,
720 embedding_sha256: Option<String>,
721 hf_token: Option<String>,
722 inference_timeout: std::time::Duration,
723}
724
725#[cfg(feature = "candle")]
726fn resolve_candle_load_params(
727 entry: &ProviderEntry,
728 candle: &CandleInlineConfig,
729 config: &Config,
730) -> CandleLoadParams {
731 let source = match candle.source {
732 CandleSource::Local => zeph_llm::candle_provider::loader::ModelSource::Local {
733 path: std::path::PathBuf::from(&candle.local_path),
734 },
735 CandleSource::Huggingface => zeph_llm::candle_provider::loader::ModelSource::HuggingFace {
736 repo_id: entry
737 .model
738 .clone()
739 .unwrap_or_else(|| config.llm.effective_model().to_owned()),
740 filename: candle.filename.clone(),
741 sha256: candle.chat_model_sha256.clone(),
742 },
743 };
744 let template =
745 zeph_llm::candle_provider::template::ChatTemplate::parse_str(&candle.chat_template);
746 let gen_config = zeph_llm::candle_provider::generate::GenerationConfig {
747 temperature: candle.generation.temperature,
748 top_p: candle.generation.top_p,
749 top_k: candle.generation.top_k,
750 max_tokens: candle.generation.capped_max_tokens(),
751 seed: candle.generation.seed,
752 repeat_penalty: candle.generation.repeat_penalty,
753 repeat_last_n: candle.generation.repeat_last_n,
754 };
755 let inference_timeout = std::time::Duration::from_secs(candle.inference_timeout_secs.max(1));
758 CandleLoadParams {
759 source,
760 template,
761 gen_config,
762 embedding_repo: candle.embedding_repo.clone(),
763 embedding_sha256: candle.embedding_model_sha256.clone(),
764 hf_token: candle.hf_token.clone(),
765 inference_timeout,
766 }
767}
768
769#[cfg(feature = "candle")]
770fn build_candle_provider(
771 entry: &ProviderEntry,
772 config: &Config,
773) -> Result<AnyProvider, BootstrapError> {
774 let candle = entry.candle.as_ref().ok_or_else(|| {
775 BootstrapError::Provider(
776 "candle provider requires 'candle' section in [[llm.providers]]".into(),
777 )
778 })?;
779 let params = resolve_candle_load_params(entry, candle, config);
780 let device = select_device(candle.device)?;
781 zeph_llm::candle_provider::CandleProvider::new_with_timeout(
782 ¶ms.source,
783 params.template,
784 params.gen_config,
785 params.embedding_repo.as_deref(),
786 params.embedding_sha256.as_deref(),
787 params.hf_token.as_deref(),
788 device,
789 params.inference_timeout,
790 )
791 .map(AnyProvider::Candle)
792 .map_err(|e| BootstrapError::Provider(e.to_string()))
793}
794
795#[cfg(feature = "candle")]
802pub fn select_device(
803 preference: CandleDevice,
804) -> Result<zeph_llm::candle_provider::Device, BootstrapError> {
805 match preference {
806 CandleDevice::Metal => {
807 #[cfg(feature = "metal")]
808 return zeph_llm::candle_provider::Device::new_metal(0)
809 .map_err(|e| BootstrapError::Provider(e.to_string()));
810 #[cfg(not(feature = "metal"))]
811 return Err(BootstrapError::Provider(
812 "candle compiled without metal feature".into(),
813 ));
814 }
815 CandleDevice::Cuda => {
816 #[cfg(feature = "cuda")]
817 return zeph_llm::candle_provider::Device::new_cuda(0)
818 .map_err(|e| BootstrapError::Provider(e.to_string()));
819 #[cfg(not(feature = "cuda"))]
820 return Err(BootstrapError::Provider(
821 "candle compiled without cuda feature".into(),
822 ));
823 }
824 CandleDevice::Cpu => Ok(zeph_llm::candle_provider::Device::Cpu),
825 CandleDevice::Auto => {
826 #[cfg(feature = "metal")]
827 if let Ok(device) = zeph_llm::candle_provider::Device::new_metal(0) {
828 return Ok(device);
829 }
830 #[cfg(feature = "cuda")]
831 if let Ok(device) = zeph_llm::candle_provider::Device::new_cuda(0) {
832 return Ok(device);
833 }
834 Ok(zeph_llm::candle_provider::Device::Cpu)
835 }
836 }
837}
838
839#[cfg(test)]
840mod tests {
841 #[cfg(feature = "candle")]
842 use super::select_device;
843 #[cfg(feature = "candle")]
844 use crate::config::CandleDevice;
845 #[cfg(feature = "candle")]
846 use std::assert_matches;
847
848 #[cfg(feature = "candle")]
849 #[test]
850 fn select_device_cpu_default() {
851 let device = select_device(CandleDevice::Cpu).unwrap();
852 assert_matches!(device, zeph_llm::candle_provider::Device::Cpu);
853 }
854
855 #[cfg(all(feature = "candle", not(feature = "metal")))]
856 #[test]
857 fn select_device_metal_without_feature_errors() {
858 let result = select_device(CandleDevice::Metal);
859 assert!(result.is_err());
860 assert!(result.unwrap_err().to_string().contains("metal feature"));
861 }
862
863 #[cfg(all(feature = "candle", not(feature = "cuda")))]
864 #[test]
865 fn select_device_cuda_without_feature_errors() {
866 let result = select_device(CandleDevice::Cuda);
867 assert!(result.is_err());
868 assert!(result.unwrap_err().to_string().contains("cuda feature"));
869 }
870
871 #[cfg(feature = "candle")]
875 #[test]
876 fn resolve_candle_load_params_threads_chat_and_embedding_sha256() {
877 use super::resolve_candle_load_params;
878 use crate::config::{CandleInlineConfig, Config};
879 use zeph_config::providers::ProviderEntry;
880
881 let candle = CandleInlineConfig {
882 chat_model_sha256: Some("deadbeef".into()),
883 embedding_repo: Some("org/embed-model".into()),
884 embedding_model_sha256: Some("cafef00d".into()),
885 ..CandleInlineConfig::default()
886 };
887 let entry = ProviderEntry {
888 model: Some("org/chat-model".into()),
889 candle: Some(candle.clone()),
890 ..ProviderEntry::default()
891 };
892 let config = Config::default();
893
894 let params = resolve_candle_load_params(&entry, &candle, &config);
895
896 if let zeph_llm::candle_provider::loader::ModelSource::HuggingFace { sha256, .. } =
897 params.source
898 {
899 assert_eq!(sha256.as_deref(), Some("deadbeef"));
900 } else {
901 panic!("expected HuggingFace source for CandleSource::default()")
902 }
903 assert_eq!(params.embedding_sha256.as_deref(), Some("cafef00d"));
904 }
905
906 #[cfg(feature = "candle")]
907 #[test]
908 fn resolve_candle_load_params_sha256_absent_by_default() {
909 use super::resolve_candle_load_params;
910 use crate::config::{CandleInlineConfig, Config};
911 use zeph_config::providers::ProviderEntry;
912
913 let candle = CandleInlineConfig::default();
914 let entry = ProviderEntry {
915 model: Some("org/chat-model".into()),
916 candle: Some(candle.clone()),
917 ..ProviderEntry::default()
918 };
919 let config = Config::default();
920
921 let params = resolve_candle_load_params(&entry, &candle, &config);
922
923 if let zeph_llm::candle_provider::loader::ModelSource::HuggingFace { sha256, .. } =
924 params.source
925 {
926 assert!(sha256.is_none());
927 } else {
928 panic!("expected HuggingFace source for CandleSource::default()")
929 }
930 assert!(params.embedding_sha256.is_none());
931 }
932
933 #[cfg(feature = "cocoon")]
934 use super::spawn_cocoon_health_checks;
935 use super::{build_provider_from_entry, resolve_named_provider};
936 use crate::config::{Config, ProviderKind};
937 use zeph_config::providers::ProviderEntry;
938 use zeph_llm::LlmProvider;
939
940 #[cfg(feature = "gonka")]
941 mod gonka_tests {
942 use super::*;
943 use zeph_common::secret::Secret;
944 use zeph_config::GonkaNode;
945 use zeph_llm::LlmProvider;
946
947 fn gonka_entry_with_nodes(nodes: Vec<GonkaNode>) -> ProviderEntry {
948 ProviderEntry {
949 provider_type: ProviderKind::Gonka,
950 name: Some("gonka".into()),
951 model: Some("gpt-4o".into()),
952 gonka_nodes: nodes,
953 ..ProviderEntry::default()
954 }
955 }
956
957 fn valid_nodes() -> Vec<GonkaNode> {
958 vec![GonkaNode {
959 url: "https://node1.gonka.ai".into(),
960 address: "gonka1w508d6qejxtdg4y5r3zarvary0c5xw7k2gsyg6".into(),
961 name: Some("node1".into()),
962 }]
963 }
964
965 const VALID_PRIV_KEY: &str =
966 "0000000000000000000000000000000000000000000000000000000000000001";
967
968 #[test]
969 fn build_gonka_provider_missing_key_returns_error() {
970 let entry = gonka_entry_with_nodes(valid_nodes());
971 let config = Config::default();
972 let result = build_provider_from_entry(&entry, &config, None);
973 assert!(result.is_err());
974 let msg = result.unwrap_err().to_string();
975 assert!(
976 msg.contains("ZEPH_GONKA_PRIVATE_KEY"),
977 "error must mention missing key: {msg}"
978 );
979 }
980
981 #[test]
982 fn build_gonka_provider_empty_nodes_returns_error() {
983 let entry = gonka_entry_with_nodes(vec![]);
984 let mut config = Config::default();
985 config.secrets.gonka_private_key = Some(Secret::new(VALID_PRIV_KEY));
986 let result = build_provider_from_entry(&entry, &config, None);
987 assert!(result.is_err());
988 let msg = result.unwrap_err().to_string();
989 assert!(
990 msg.contains("gonka_nodes") || msg.contains("node"),
991 "error must mention empty nodes: {msg}"
992 );
993 }
994
995 #[test]
996 fn build_gonka_provider_address_mismatch_returns_error() {
997 let entry = gonka_entry_with_nodes(valid_nodes());
998 let mut config = Config::default();
999 config.secrets.gonka_private_key = Some(Secret::new(VALID_PRIV_KEY));
1000 config.secrets.gonka_address =
1001 Some(Secret::new("gonka1wrongaddress000000000000000000000000000"));
1002 let result = build_provider_from_entry(&entry, &config, None);
1003 assert!(result.is_err());
1004 let msg = result.unwrap_err().to_string();
1005 assert!(
1006 msg.contains("does not match"),
1007 "error must mention address mismatch: {msg}"
1008 );
1009 }
1010
1011 #[test]
1012 fn build_gonka_provider_happy_path() {
1013 let entry = gonka_entry_with_nodes(valid_nodes());
1014 let mut config = Config::default();
1015 config.secrets.gonka_private_key = Some(Secret::new(VALID_PRIV_KEY));
1016 let result = build_provider_from_entry(&entry, &config, None);
1017 assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
1018 let provider = result.unwrap();
1019 assert_eq!(provider.name(), "gonka");
1020 }
1021 }
1022
1023 fn make_provider_entry(
1024 embed: bool,
1025 model: Option<&str>,
1026 embedding_model: Option<&str>,
1027 ) -> ProviderEntry {
1028 ProviderEntry {
1029 provider_type: ProviderKind::Ollama,
1030 embed,
1031 model: model.map(str::to_owned),
1032 embedding_model: embedding_model.map(str::to_owned),
1033 ..ProviderEntry::default()
1034 }
1035 }
1036
1037 #[test]
1038 fn stable_skill_embedding_model_prefers_embedding_model_field() {
1039 let mut config = Config::default();
1040 config.llm.providers = vec![make_provider_entry(
1041 true,
1042 Some("chat-model"),
1043 Some("embed-v2"),
1044 )];
1045 assert_eq!(config.llm.stable_skill_embedding_model(), "embed-v2");
1046 }
1047
1048 #[test]
1049 fn stable_skill_embedding_model_falls_back_to_model_field() {
1050 let mut config = Config::default();
1051 config.llm.providers = vec![make_provider_entry(
1052 true,
1053 Some("nomic-embed-text-v2-moe:latest"),
1054 None,
1055 )];
1056 assert_eq!(
1057 config.llm.stable_skill_embedding_model(),
1058 "nomic-embed-text-v2-moe:latest"
1059 );
1060 }
1061
1062 #[test]
1063 fn stable_skill_embedding_model_finds_embed_flag_entry() {
1064 let mut config = Config::default();
1065 config.llm.providers = vec![
1066 make_provider_entry(false, Some("chat-model"), None),
1067 make_provider_entry(true, Some("embed-model"), Some("text-embed-3")),
1068 ];
1069 assert_eq!(config.llm.stable_skill_embedding_model(), "text-embed-3");
1070 }
1071
1072 #[test]
1073 fn stable_skill_embedding_model_falls_back_to_effective_when_no_embed_entry() {
1074 let mut config = Config::default();
1075 config.llm.embedding_model = "global-embed-model".to_owned();
1076 config.llm.providers = vec![make_provider_entry(false, Some("chat"), None)];
1078 assert_eq!(
1079 config.llm.stable_skill_embedding_model(),
1080 config.llm.effective_embedding_model()
1081 );
1082 }
1083
1084 #[test]
1085 fn resolve_named_provider_empty_name_falls_back_to_primary_silently() {
1086 let config = Config::default();
1087 let entry = make_provider_entry(false, Some("primary-model"), None);
1088 let primary = build_provider_from_entry(&entry, &config, None).unwrap();
1089 let resolved = resolve_named_provider(&config, &primary, "");
1090 assert_eq!(resolved.name(), primary.name());
1091 }
1092
1093 #[test]
1094 fn resolve_named_provider_unmatched_name_falls_back_to_primary_with_warn() {
1095 let config = Config::default(); let entry = make_provider_entry(false, Some("primary-model"), None);
1097 let primary = build_provider_from_entry(&entry, &config, None).unwrap();
1098 let resolved = resolve_named_provider(&config, &primary, "fast");
1099 assert_eq!(resolved.name(), primary.name());
1100 }
1101
1102 #[cfg(feature = "cocoon")]
1103 mod cocoon_tests {
1104 use super::*;
1105
1106 fn cocoon_entry(access_hash: Option<&str>) -> ProviderEntry {
1107 ProviderEntry {
1108 provider_type: ProviderKind::Cocoon,
1109 name: Some("cocoon".into()),
1110 model: Some("Qwen/Qwen3-0.6B".into()),
1111 cocoon_client_url: Some("http://localhost:10000".into()),
1112 cocoon_access_hash: access_hash.map(str::to_owned),
1113 cocoon_health_check: false,
1114 ..ProviderEntry::default()
1115 }
1116 }
1117
1118 #[test]
1120 fn cocoon_access_hash_gate_vault_miss_errors() {
1121 let entry = cocoon_entry(Some(""));
1122 let config = Config::default(); let result = build_provider_from_entry(&entry, &config, None);
1124 assert!(
1125 result.is_err(),
1126 "expected error when vault key is absent but sentinel is set"
1127 );
1128 let err_str = result.unwrap_err().to_string();
1129 assert!(
1130 err_str.contains("ZEPH_COCOON_ACCESS_HASH"),
1131 "error should mention the vault key: {err_str}"
1132 );
1133 }
1134
1135 #[test]
1137 fn cocoon_no_access_hash_gate_succeeds_without_vault() {
1138 let entry = cocoon_entry(None);
1139 let config = Config::default();
1140 let result = build_provider_from_entry(&entry, &config, None);
1141 assert!(
1142 result.is_ok(),
1143 "expected success when no access hash requested: {:?}",
1144 result.err()
1145 );
1146 }
1147
1148 #[test]
1152 fn spawn_cocoon_health_checks_skips_on_vault_miss() {
1153 let mut entry = cocoon_entry(Some(""));
1154 entry.cocoon_health_check = true;
1155 let config = Config::default(); let supervisor = std::sync::Arc::new(zeph_common::TaskSupervisor::new(
1157 tokio_util::sync::CancellationToken::new(),
1158 ));
1159
1160 spawn_cocoon_health_checks(&[&entry], &config, &supervisor);
1161
1162 assert!(
1163 supervisor.snapshot().is_empty(),
1164 "expected no health-check task to be spawned when vault key is absent"
1165 );
1166 }
1167 }
1168}