Skip to main content

ferrum_engine/
registry.rs

1//! Component registry for dynamic component creation
2//!
3//! This module provides a registry pattern implementation that allows
4//! dynamic registration and lookup of component factories. The registry
5//! replaces hardcoded factory implementations with a flexible, extensible
6//! system that supports:
7//!
8//! - Dynamic component registration and discovery
9//! - Multiple implementations for each component type
10//! - Configuration-driven component selection
11//! - Easy testing with mock components
12
13use async_trait::async_trait;
14use ferrum_interfaces::{
15    KvCacheManager, ModelExecutor, Sampler, SchedulerInterface as Scheduler, Tokenizer,
16};
17use ferrum_models::vnext::{PreparedProductionModel, ProductionModelSourceBundle};
18use ferrum_types::{Device, EngineConfig, FerrumError, Result, RuntimeKnobs};
19use parking_lot::RwLock;
20use std::collections::HashMap;
21use std::sync::{Arc, OnceLock};
22use tracing::{debug, info};
23
24pub use ferrum_interfaces::sampler::GreedySampler;
25
26// ============================================================================
27// Core Types
28// ============================================================================
29
30/// Factory trait for creating components
31#[async_trait]
32pub trait ComponentFactory<T>: Send + Sync {
33    /// Create a component instance
34    async fn create(&self, config: &ComponentConfig) -> Result<T>;
35
36    /// Get component metadata
37    fn metadata(&self) -> ComponentMetadata;
38}
39
40/// Component configuration passed to factories
41#[derive(Debug, Clone)]
42pub struct ComponentConfig {
43    /// Full engine configuration
44    pub engine_config: EngineConfig,
45    /// Target device for the component
46    pub device: Device,
47    /// Additional component-specific options
48    pub component_options: HashMap<String, serde_json::Value>,
49    /// Immutable role-specific product sources. This is deliberately absent
50    /// from serialized EngineConfig and shared by tokenizer and executor.
51    pub model_sources: Option<Arc<ProductionModelSourceBundle>>,
52    /// Prepared typed model derived from `model_sources`, when the composition
53    /// root has already resolved a migrated family.
54    pub prepared_model: Option<Arc<PreparedProductionModel>>,
55}
56
57impl ComponentConfig {
58    /// Create from engine config
59    pub fn from_engine_config(config: &EngineConfig) -> Self {
60        Self::from_engine_config_and_product_model(config, None, None)
61    }
62
63    pub fn from_engine_config_and_sources(
64        config: &EngineConfig,
65        model_sources: Option<Arc<ProductionModelSourceBundle>>,
66    ) -> Self {
67        Self::from_engine_config_and_product_model(config, model_sources, None)
68    }
69
70    pub fn from_engine_config_and_product_model(
71        config: &EngineConfig,
72        model_sources: Option<Arc<ProductionModelSourceBundle>>,
73        prepared_model: Option<Arc<PreparedProductionModel>>,
74    ) -> Self {
75        Self {
76            engine_config: config.clone(),
77            device: config.backend.device.clone(),
78            component_options: config.backend.backend_options.clone(),
79            model_sources,
80            prepared_model,
81        }
82    }
83
84    /// Get an option value
85    pub fn get_option<T: serde::de::DeserializeOwned>(&self, key: &str) -> Option<T> {
86        self.component_options
87            .get(key)
88            .and_then(|v| serde_json::from_value(v.clone()).ok())
89    }
90
91    /// Get string option
92    pub fn get_string_option(&self, key: &str) -> Option<String> {
93        self.component_options
94            .get(key)
95            .and_then(|v| v.as_str())
96            .map(|s| s.to_string())
97    }
98}
99
100/// Component metadata
101#[derive(Debug, Clone)]
102pub struct ComponentMetadata {
103    /// Component name
104    pub name: String,
105    /// Version string
106    pub version: String,
107    /// Human-readable description
108    pub description: String,
109    /// Supported devices
110    pub supported_devices: Vec<Device>,
111    /// Capability flags
112    pub capabilities: Vec<String>,
113}
114
115impl Default for ComponentMetadata {
116    fn default() -> Self {
117        Self {
118            name: "unknown".to_string(),
119            version: "0.0.0".to_string(),
120            description: String::new(),
121            supported_devices: vec![Device::CPU],
122            capabilities: vec![],
123        }
124    }
125}
126
127fn cpu_cuda_and_optional_metal_devices() -> Vec<Device> {
128    #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
129    {
130        vec![Device::CPU, Device::CUDA(0), Device::Metal]
131    }
132    #[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
133    {
134        vec![Device::CPU, Device::CUDA(0)]
135    }
136}
137
138fn parse_executor_dtype(s: &str) -> Option<candle_core::DType> {
139    match s.trim().to_ascii_lowercase().as_str() {
140        "fp16" | "f16" | "float16" => Some(candle_core::DType::F16),
141        "fp32" | "f32" | "float32" => Some(candle_core::DType::F32),
142        _ => None,
143    }
144}
145
146#[derive(Debug, Clone, PartialEq, Eq)]
147struct RegistryRuntimeEnv {
148    model_path: Option<String>,
149    metal_dtype: Option<String>,
150    dtype: Option<String>,
151    tp: usize,
152}
153
154impl RegistryRuntimeEnv {
155    /// Build from the typed runtime snapshot resolved at the composition root.
156    /// Replaces the former `std::env`-reading `from_env`/`OnceLock` pair; the
157    /// CLI lands FERRUM_MODEL_PATH / FERRUM_DTYPE / FERRUM_METAL_DTYPE /
158    /// FERRUM_TP into `EngineConfig.runtime`, and the registry reads that.
159    fn from_runtime_knobs(knobs: &RuntimeKnobs) -> Self {
160        Self {
161            model_path: knobs.model_path.clone(),
162            metal_dtype: knobs.metal_dtype.clone(),
163            dtype: knobs.dtype.clone(),
164            tp: knobs.tp.unwrap_or(0),
165        }
166    }
167
168    #[cfg(test)]
169    fn from_env_vars<I, K, V>(vars: I) -> Self
170    where
171        I: IntoIterator<Item = (K, V)>,
172        K: AsRef<str>,
173        V: Into<String>,
174    {
175        let mut model_path = None;
176        let mut metal_dtype = None;
177        let mut dtype = None;
178        let mut tp = None;
179
180        for (key, value) in vars {
181            let value = value.into();
182            match key.as_ref() {
183                "FERRUM_MODEL_PATH" => model_path = Some(value),
184                "FERRUM_METAL_DTYPE" => metal_dtype = Some(value),
185                "FERRUM_DTYPE" => dtype = Some(value),
186                "FERRUM_TP" => tp = value.parse::<usize>().ok(),
187                _ => {}
188            }
189        }
190
191        Self {
192            model_path,
193            metal_dtype,
194            dtype,
195            tp: tp.unwrap_or(0),
196        }
197    }
198
199    #[cfg(test)]
200    fn model_path(&self) -> Option<String> {
201        self.model_path.clone()
202    }
203
204    fn dtype_for_device(&self, device: &Device) -> candle_core::DType {
205        match device {
206            Device::CPU => candle_core::DType::F32,
207            #[cfg(any(target_os = "macos", target_os = "ios"))]
208            Device::Metal => self
209                .metal_dtype
210                .as_deref()
211                .and_then(parse_executor_dtype)
212                .or_else(|| self.dtype.as_deref().and_then(parse_executor_dtype))
213                .unwrap_or(candle_core::DType::F32),
214            Device::CUDA(_) | Device::ROCm(_) => self
215                .dtype
216                .as_deref()
217                .and_then(parse_executor_dtype)
218                .unwrap_or(candle_core::DType::F16),
219        }
220    }
221
222    #[allow(dead_code)]
223    fn metal_dtype_hint(&self) -> String {
224        self.metal_dtype
225            .clone()
226            .unwrap_or_else(|| "f32".to_string())
227    }
228}
229
230// ============================================================================
231// Component Registry
232// ============================================================================
233
234/// Global component registry for managing component factories
235pub struct ComponentRegistry {
236    tokenizer_factories:
237        RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>>>,
238    sampler_factories:
239        RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>>>,
240    scheduler_factories:
241        RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>>>,
242    kv_cache_factories:
243        RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>>>,
244    executor_factories:
245        RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>>>,
246}
247
248impl ComponentRegistry {
249    /// Create a new empty registry
250    pub fn new() -> Self {
251        Self {
252            tokenizer_factories: RwLock::new(HashMap::new()),
253            sampler_factories: RwLock::new(HashMap::new()),
254            scheduler_factories: RwLock::new(HashMap::new()),
255            kv_cache_factories: RwLock::new(HashMap::new()),
256            executor_factories: RwLock::new(HashMap::new()),
257        }
258    }
259
260    /// Create a registry with all default factories pre-registered
261    pub fn with_defaults() -> Self {
262        let registry = Self::new();
263        registry.register_defaults();
264        registry
265    }
266
267    /// Register all default component factories
268    pub fn register_defaults(&self) {
269        info!("Registering default component factories");
270
271        // Tokenizer factories
272        self.register_tokenizer_factory("huggingface", Arc::new(HuggingFaceTokenizerFactory));
273        self.register_tokenizer_factory("stub", Arc::new(StubTokenizerFactory));
274
275        // Sampler factories
276        self.register_sampler_factory("multinomial", Arc::new(MultinomialSamplerFactory));
277        self.register_sampler_factory("greedy", Arc::new(GreedySamplerFactory));
278
279        // Scheduler factories
280        self.register_scheduler_factory("fifo", Arc::new(FifoSchedulerFactory));
281        self.register_scheduler_factory("priority", Arc::new(PrioritySchedulerFactory));
282        self.register_scheduler_factory("continuous", Arc::new(ContinuousBatchSchedulerFactory));
283
284        // KV cache factories
285        self.register_kv_cache_factory("default", Arc::new(DefaultKvCacheFactory));
286        self.register_kv_cache_factory("paged", Arc::new(PagedKvCacheFactory));
287
288        // Executor factories
289        self.register_executor_factory("stub", Arc::new(StubExecutorFactory));
290        self.register_executor_factory("llm", Arc::new(LlmExecutorFactory));
291
292        debug!(
293            "Registered factories - tokenizers: {:?}, samplers: {:?}, schedulers: {:?}, kv_caches: {:?}, executors: {:?}",
294            self.list_tokenizers(),
295            self.list_samplers(),
296            self.list_schedulers(),
297            self.list_kv_caches(),
298            self.list_executors()
299        );
300    }
301
302    // ========================================================================
303    // Registration methods
304    // ========================================================================
305
306    /// Register a tokenizer factory
307    pub fn register_tokenizer_factory(
308        &self,
309        name: impl Into<String>,
310        factory: Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>,
311    ) {
312        let name = name.into();
313        debug!("Registering tokenizer factory: {}", name);
314        self.tokenizer_factories.write().insert(name, factory);
315    }
316
317    /// Register a sampler factory
318    pub fn register_sampler_factory(
319        &self,
320        name: impl Into<String>,
321        factory: Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>,
322    ) {
323        let name = name.into();
324        debug!("Registering sampler factory: {}", name);
325        self.sampler_factories.write().insert(name, factory);
326    }
327
328    /// Register a scheduler factory
329    pub fn register_scheduler_factory(
330        &self,
331        name: impl Into<String>,
332        factory: Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>,
333    ) {
334        let name = name.into();
335        debug!("Registering scheduler factory: {}", name);
336        self.scheduler_factories.write().insert(name, factory);
337    }
338
339    /// Register a KV cache factory
340    pub fn register_kv_cache_factory(
341        &self,
342        name: impl Into<String>,
343        factory: Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>,
344    ) {
345        let name = name.into();
346        debug!("Registering KV cache factory: {}", name);
347        self.kv_cache_factories.write().insert(name, factory);
348    }
349
350    /// Register a model executor factory
351    pub fn register_executor_factory(
352        &self,
353        name: impl Into<String>,
354        factory: Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>,
355    ) {
356        let name = name.into();
357        debug!("Registering executor factory: {}", name);
358        self.executor_factories.write().insert(name, factory);
359    }
360
361    // ========================================================================
362    // Lookup methods
363    // ========================================================================
364
365    /// Get a tokenizer factory by name
366    pub fn get_tokenizer_factory(
367        &self,
368        name: &str,
369    ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>> {
370        self.tokenizer_factories.read().get(name).cloned()
371    }
372
373    /// Get a sampler factory by name
374    pub fn get_sampler_factory(
375        &self,
376        name: &str,
377    ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>> {
378        self.sampler_factories.read().get(name).cloned()
379    }
380
381    /// Get a scheduler factory by name
382    pub fn get_scheduler_factory(
383        &self,
384        name: &str,
385    ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>> {
386        self.scheduler_factories.read().get(name).cloned()
387    }
388
389    /// Get a KV cache factory by name
390    pub fn get_kv_cache_factory(
391        &self,
392        name: &str,
393    ) -> Option<Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>> {
394        self.kv_cache_factories.read().get(name).cloned()
395    }
396
397    /// Get a model executor factory by name
398    pub fn get_executor_factory(
399        &self,
400        name: &str,
401    ) -> Option<Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>> {
402        self.executor_factories.read().get(name).cloned()
403    }
404
405    // ========================================================================
406    // List methods
407    // ========================================================================
408
409    /// List all registered tokenizer names
410    pub fn list_tokenizers(&self) -> Vec<String> {
411        self.tokenizer_factories.read().keys().cloned().collect()
412    }
413
414    /// List all registered sampler names
415    pub fn list_samplers(&self) -> Vec<String> {
416        self.sampler_factories.read().keys().cloned().collect()
417    }
418
419    /// List all registered scheduler names
420    pub fn list_schedulers(&self) -> Vec<String> {
421        self.scheduler_factories.read().keys().cloned().collect()
422    }
423
424    /// List all registered KV cache names
425    pub fn list_kv_caches(&self) -> Vec<String> {
426        self.kv_cache_factories.read().keys().cloned().collect()
427    }
428
429    /// List all registered executor names
430    pub fn list_executors(&self) -> Vec<String> {
431        self.executor_factories.read().keys().cloned().collect()
432    }
433
434    // ========================================================================
435    // Convenience creation methods
436    // ========================================================================
437
438    /// Create a tokenizer by name
439    pub async fn create_tokenizer(
440        &self,
441        name: &str,
442        config: &ComponentConfig,
443    ) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
444        let factory = self.get_tokenizer_factory(name).ok_or_else(|| {
445            FerrumError::tokenizer(format!(
446                "Tokenizer '{}' not found. Available: {:?}",
447                name,
448                self.list_tokenizers()
449            ))
450        })?;
451        factory.create(config).await
452    }
453
454    /// Create a sampler by name
455    pub async fn create_sampler(
456        &self,
457        name: &str,
458        config: &ComponentConfig,
459    ) -> Result<Arc<dyn Sampler + Send + Sync>> {
460        let factory = self.get_sampler_factory(name).ok_or_else(|| {
461            FerrumError::internal(format!(
462                "Sampler '{}' not found. Available: {:?}",
463                name,
464                self.list_samplers()
465            ))
466        })?;
467        factory.create(config).await
468    }
469
470    /// Create a scheduler by name
471    pub async fn create_scheduler(
472        &self,
473        name: &str,
474        config: &ComponentConfig,
475    ) -> Result<Arc<dyn Scheduler + Send + Sync>> {
476        let factory = self.get_scheduler_factory(name).ok_or_else(|| {
477            FerrumError::scheduler(format!(
478                "Scheduler '{}' not found. Available: {:?}",
479                name,
480                self.list_schedulers()
481            ))
482        })?;
483        factory.create(config).await
484    }
485
486    /// Create a KV cache manager by name
487    pub async fn create_kv_cache(
488        &self,
489        name: &str,
490        config: &ComponentConfig,
491    ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
492        let factory = self.get_kv_cache_factory(name).ok_or_else(|| {
493            FerrumError::internal(format!(
494                "KV cache '{}' not found. Available: {:?}",
495                name,
496                self.list_kv_caches()
497            ))
498        })?;
499        factory.create(config).await
500    }
501
502    /// Create a model executor by name
503    pub async fn create_executor(
504        &self,
505        name: &str,
506        config: &ComponentConfig,
507    ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
508        let factory = self.get_executor_factory(name).ok_or_else(|| {
509            FerrumError::model(format!(
510                "Executor '{}' not found. Available: {:?}",
511                name,
512                self.list_executors()
513            ))
514        })?;
515        factory.create(config).await
516    }
517}
518
519impl Default for ComponentRegistry {
520    fn default() -> Self {
521        Self::with_defaults()
522    }
523}
524
525impl std::fmt::Debug for ComponentRegistry {
526    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
527        f.debug_struct("ComponentRegistry")
528            .field("tokenizers", &self.list_tokenizers())
529            .field("samplers", &self.list_samplers())
530            .field("schedulers", &self.list_schedulers())
531            .field("kv_caches", &self.list_kv_caches())
532            .field("executors", &self.list_executors())
533            .finish()
534    }
535}
536
537// ============================================================================
538// Default Factory Implementations
539// ============================================================================
540
541// ----------------------------------------------------------------------------
542// Backend Factories
543// ----------------------------------------------------------------------------
544
545// CandleBackendFactory deleted in Phase: legacy `ComputeBackend` trait
546// gone; the registry's component graph never produced a runtime backend
547// (`_backend` in builder.rs was unused). The stub-executor factory now
548// constructs `CandleTensorFactory` directly when it needs to mint dummy
549// logits; everything else dispatches via `Backend<B>` in ferrum-kernels.
550
551// ----------------------------------------------------------------------------
552// Tokenizer Factories
553// ----------------------------------------------------------------------------
554
555/// HuggingFace tokenizer factory
556pub struct HuggingFaceTokenizerFactory;
557
558#[async_trait]
559impl ComponentFactory<Arc<dyn Tokenizer + Send + Sync>> for HuggingFaceTokenizerFactory {
560    async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
561        if let Some(sources) = config.model_sources.as_ref() {
562            let tokenizer =
563                ferrum_tokenizer::implementations::HuggingFaceTokenizer::from_source_bytes(
564                    sources.tokenizer_json(),
565                    sources.tokenizer_config_json(),
566                    sources.generation_config_json(),
567                )
568                .await?;
569            info!(
570                tokenizer_root = %sources.tokenizer_root().display(),
571                "Loaded HuggingFace tokenizer from typed product sources"
572            );
573            return Ok(Arc::new(tokenizer));
574        }
575
576        // Try to find tokenizer path from config or environment
577        let tokenizer_path = config
578            .get_string_option("tokenizer_path")
579            .or_else(|| config.get_string_option("model_path"))
580            .or_else(|| config.engine_config.runtime.model_path.clone());
581
582        if let Some(model_path) = tokenizer_path {
583            let path = std::path::Path::new(&model_path);
584            // GGUF path: model_path is a file. Auto-discover a sibling
585            // tokenizer.json (or the convention used by ~/ferrum-bench).
586            let tokenizer_file = if ferrum_models::gguf_engine_loader::is_gguf_path(&model_path) {
587                ferrum_models::gguf_engine_loader::auto_discover_tokenizer_path(path).ok_or_else(
588                    || {
589                        FerrumError::tokenizer(format!(
590                            "Could not find tokenizer.json for {} — \
591                             place it next to the .gguf file or in a \
592                             sibling tokenizers/ directory",
593                            path.display()
594                        ))
595                    },
596                )?
597            } else {
598                // HF safetensors layout: model_path is a directory.
599                path.join("tokenizer.json")
600            };
601
602            if tokenizer_file.exists() {
603                info!("Loading HuggingFace tokenizer from: {:?}", tokenizer_file);
604                match ferrum_tokenizer::implementations::HuggingFaceTokenizer::from_file(
605                    &tokenizer_file.to_string_lossy(),
606                )
607                .await
608                {
609                    Ok(tokenizer) => {
610                        info!("HuggingFace tokenizer loaded successfully");
611                        return Ok(Arc::new(tokenizer));
612                    }
613                    Err(e) => {
614                        tracing::warn!("Failed to load tokenizer: {}, falling back to stub", e);
615                    }
616                }
617            }
618        }
619
620        // Fallback to stub
621        Err(FerrumError::tokenizer(
622            "HuggingFace tokenizer path not found or invalid",
623        ))
624    }
625
626    fn metadata(&self) -> ComponentMetadata {
627        ComponentMetadata {
628            name: "huggingface".to_string(),
629            version: "0.1.0".to_string(),
630            description: "HuggingFace tokenizers library integration".to_string(),
631            supported_devices: vec![Device::CPU],
632            capabilities: vec![
633                "bpe".to_string(),
634                "wordpiece".to_string(),
635                "sentencepiece".to_string(),
636                "chat_template".to_string(),
637            ],
638        }
639    }
640}
641
642/// Stub tokenizer factory for testing
643pub struct StubTokenizerFactory;
644
645#[async_trait]
646impl ComponentFactory<Arc<dyn Tokenizer + Send + Sync>> for StubTokenizerFactory {
647    async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
648        let vocab_size = config
649            .engine_config
650            .model
651            .model_info
652            .as_ref()
653            .map(|info| info.vocab_size)
654            .unwrap_or(32000);
655
656        info!("Creating stub tokenizer with vocab_size: {}", vocab_size);
657        Ok(Arc::new(StubTokenizer::new(vocab_size)))
658    }
659
660    fn metadata(&self) -> ComponentMetadata {
661        ComponentMetadata {
662            name: "stub".to_string(),
663            version: "0.1.0".to_string(),
664            description: "Stub tokenizer for testing".to_string(),
665            supported_devices: vec![Device::CPU],
666            capabilities: vec!["testing".to_string()],
667        }
668    }
669}
670
671// ----------------------------------------------------------------------------
672// Sampler Factories
673// ----------------------------------------------------------------------------
674
675/// Multinomial sampler factory
676pub struct MultinomialSamplerFactory;
677
678#[async_trait]
679impl ComponentFactory<Arc<dyn Sampler + Send + Sync>> for MultinomialSamplerFactory {
680    async fn create(&self, _config: &ComponentConfig) -> Result<Arc<dyn Sampler + Send + Sync>> {
681        info!("Creating multinomial sampler");
682        Ok(Arc::new(ferrum_interfaces::sampler::MultinomialSampler))
683    }
684
685    fn metadata(&self) -> ComponentMetadata {
686        ComponentMetadata {
687            name: "multinomial".to_string(),
688            version: "0.1.0".to_string(),
689            description: "Multinomial sampling with temperature and top-k/top-p".to_string(),
690            supported_devices: vec![Device::CPU],
691            capabilities: vec![
692                "temperature".to_string(),
693                "top_k".to_string(),
694                "top_p".to_string(),
695            ],
696        }
697    }
698}
699
700/// Greedy sampler factory
701pub struct GreedySamplerFactory;
702
703#[async_trait]
704impl ComponentFactory<Arc<dyn Sampler + Send + Sync>> for GreedySamplerFactory {
705    async fn create(&self, _config: &ComponentConfig) -> Result<Arc<dyn Sampler + Send + Sync>> {
706        info!("Creating greedy sampler");
707        Ok(Arc::new(GreedySampler))
708    }
709
710    fn metadata(&self) -> ComponentMetadata {
711        ComponentMetadata {
712            name: "greedy".to_string(),
713            version: "0.1.0".to_string(),
714            description: "Greedy decoding (always pick highest probability)".to_string(),
715            supported_devices: vec![Device::CPU],
716            capabilities: vec!["deterministic".to_string()],
717        }
718    }
719}
720
721// ----------------------------------------------------------------------------
722// Scheduler Factories
723// ----------------------------------------------------------------------------
724
725/// FIFO scheduler factory
726pub struct FifoSchedulerFactory;
727
728#[async_trait]
729impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for FifoSchedulerFactory {
730    async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
731        info!("Creating FIFO scheduler");
732        let scheduler_config = config.engine_config.scheduler.clone();
733        let scheduler = ferrum_scheduler::implementations::FifoScheduler::new(scheduler_config);
734        Ok(Arc::new(scheduler))
735    }
736
737    fn metadata(&self) -> ComponentMetadata {
738        ComponentMetadata {
739            name: "fifo".to_string(),
740            version: "0.1.0".to_string(),
741            description: "First-In-First-Out scheduler".to_string(),
742            supported_devices: vec![Device::CPU],
743            capabilities: vec!["simple".to_string(), "fair".to_string()],
744        }
745    }
746}
747
748/// Priority scheduler factory
749pub struct PrioritySchedulerFactory;
750
751#[async_trait]
752impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for PrioritySchedulerFactory {
753    async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
754        info!("Creating priority scheduler");
755        let scheduler_config = config.engine_config.scheduler.clone();
756        let scheduler = ferrum_scheduler::implementations::PriorityScheduler::new(scheduler_config);
757        Ok(Arc::new(scheduler))
758    }
759
760    fn metadata(&self) -> ComponentMetadata {
761        ComponentMetadata {
762            name: "priority".to_string(),
763            version: "0.1.0".to_string(),
764            description: "Priority-based scheduler".to_string(),
765            supported_devices: vec![Device::CPU],
766            capabilities: vec!["priority".to_string(), "preemption".to_string()],
767        }
768    }
769}
770
771/// Continuous batching scheduler factory
772pub struct ContinuousBatchSchedulerFactory;
773
774#[async_trait]
775impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for ContinuousBatchSchedulerFactory {
776    async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
777        info!("Creating continuous batch scheduler");
778        let scheduler_config = config.engine_config.scheduler.clone();
779        let scheduler =
780            ferrum_scheduler::implementations::ContinuousBatchScheduler::new(scheduler_config);
781        Ok(Arc::new(scheduler))
782    }
783
784    fn metadata(&self) -> ComponentMetadata {
785        ComponentMetadata {
786            name: "continuous".to_string(),
787            version: "0.1.0".to_string(),
788            description: "Continuous batching scheduler with iteration-level scheduling"
789                .to_string(),
790            supported_devices: cpu_cuda_and_optional_metal_devices(),
791            capabilities: vec![
792                "continuous_batching".to_string(),
793                "preemption".to_string(),
794                "chunked_prefill".to_string(),
795                "iteration_level".to_string(),
796            ],
797        }
798    }
799}
800
801// ----------------------------------------------------------------------------
802// KV Cache Factories
803// ----------------------------------------------------------------------------
804
805/// Default KV cache factory
806pub struct DefaultKvCacheFactory;
807
808#[async_trait]
809impl ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>> for DefaultKvCacheFactory {
810    async fn create(
811        &self,
812        config: &ComponentConfig,
813    ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
814        let block_size = config.engine_config.kv_cache.block_size;
815        let max_blocks = config.engine_config.kv_cache.max_blocks;
816
817        info!(
818            "Creating default KV cache manager: device={:?}, block_size={}, max_blocks={}",
819            config.device, block_size, max_blocks
820        );
821
822        let manager = ferrum_kv::managers::DefaultKvCacheManager::new(
823            config.device.clone(),
824            block_size,
825            max_blocks,
826        )?;
827        Ok(Arc::new(manager))
828    }
829
830    fn metadata(&self) -> ComponentMetadata {
831        ComponentMetadata {
832            name: "default".to_string(),
833            version: "0.1.0".to_string(),
834            description: "Default contiguous KV cache manager".to_string(),
835            supported_devices: cpu_cuda_and_optional_metal_devices(),
836            capabilities: vec!["contiguous".to_string()],
837        }
838    }
839}
840
841/// Paged KV cache factory (for PagedAttention)
842pub struct PagedKvCacheFactory;
843
844#[async_trait]
845impl ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>> for PagedKvCacheFactory {
846    async fn create(
847        &self,
848        config: &ComponentConfig,
849    ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
850        let block_size = config.engine_config.kv_cache.block_size;
851        let max_blocks = config.engine_config.kv_cache.max_blocks;
852
853        info!(
854            "Creating paged KV cache manager: device={:?}, block_size={}, max_blocks={}",
855            config.device, block_size, max_blocks
856        );
857
858        // Use the PagedKvCacheManager for PagedAttention support
859        let paged_config = ferrum_kv::managers::PagedKvCacheConfig {
860            block_size,
861            max_gpu_blocks: max_blocks,
862            max_cpu_blocks: max_blocks / 2,
863            enable_cow: true,
864            enable_swapping: true,
865            ..Default::default()
866        };
867
868        let manager =
869            ferrum_kv::managers::PagedKvCacheManager::new(config.device.clone(), paged_config)?;
870        Ok(Arc::new(manager))
871    }
872
873    fn metadata(&self) -> ComponentMetadata {
874        ComponentMetadata {
875            name: "paged".to_string(),
876            version: "0.1.0".to_string(),
877            description: "Paged KV cache manager for PagedAttention".to_string(),
878            supported_devices: cpu_cuda_and_optional_metal_devices(),
879            capabilities: vec![
880                "paged".to_string(),
881                "copy_on_write".to_string(),
882                "swap".to_string(),
883            ],
884        }
885    }
886}
887
888// ----------------------------------------------------------------------------
889// Executor Factories
890// ----------------------------------------------------------------------------
891
892/// Stub executor factory
893pub struct StubExecutorFactory;
894
895#[async_trait]
896impl ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>> for StubExecutorFactory {
897    async fn create(
898        &self,
899        config: &ComponentConfig,
900    ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
901        let vocab_size = config
902            .engine_config
903            .model
904            .model_info
905            .as_ref()
906            .map(|info| info.vocab_size)
907            .unwrap_or(32000);
908
909        info!(
910            "Creating stub executor for model: {}",
911            config.engine_config.model.model_id
912        );
913
914        // The stub executor only needs a `TensorFactory` to mint dummy
915        // logits — no real backend dispatch involved. Wire the candle
916        // tensor factory directly without the legacy `ComputeBackend`
917        // wrapper.
918        let tensor_factory: Arc<dyn ferrum_interfaces::TensorFactory> = Arc::new(
919            crate::tensor_factory::candle::CandleTensorFactory::new(config.device.clone()),
920        );
921
922        let executor = ferrum_models::StubModelExecutor::new(
923            config.engine_config.model.model_id.clone(),
924            vocab_size,
925            tensor_factory,
926        );
927
928        Ok(Arc::new(executor))
929    }
930
931    fn metadata(&self) -> ComponentMetadata {
932        ComponentMetadata {
933            name: "stub".to_string(),
934            version: "0.1.0".to_string(),
935            description: "Stub executor for testing".to_string(),
936            supported_devices: vec![Device::CPU],
937            capabilities: vec!["testing".to_string()],
938        }
939    }
940}
941
942/// Candle executor factory
943/// Builds a [`ModelExecutor`] from a resolved model path. Despite the
944/// historical "candle" name (PR #127 deleted the legacy `CandleBackend`
945/// trait), this factory now produces real `Backend<B>`-based executors
946/// for LLMs (LlamaFamilyModel / Qwen3MoeModel) and candle-based
947/// executors for embedding / multimodal / ASR (Bert, CLIP, Whisper).
948///
949/// The factory dispatches over four axes:
950///   - Dim 1 (architecture): LlamaFamilyModel vs Qwen3MoeModel — runtime
951///     value, decided inside [`build_llm`].
952///   - Dim 3 (weight format): safetensors vs GGUF — runtime branch on
953///     `WeightFormat::detect()` at the top of [`Self::create`].
954///   - Dim 4 (device): CpuBackend / MetalBackend / CudaBackend — type
955///     parameter `B`, picked by the `(device, kv_dtype)` cascade below.
956///   - Dim 5 (kv dtype): KvFp16 / KvInt8 / KvFp8 — type parameter `K`,
957///     same cascade. Only `KvFp16` is wired today; `KvInt8` lands in
958///     PR C (see `docs/dim5-model-wireup-plan.md`).
959pub struct LlmExecutorFactory;
960
961fn resolve_llama_layer_split_plan(
962    config: &ComponentConfig,
963    num_layers: usize,
964) -> Result<Option<crate::layer_split::ParsedLayerSplitPlan>> {
965    if config
966        .get_string_option("selected_distributed_strategy")
967        .as_deref()
968        != Some("layer_split")
969    {
970        return Ok(None);
971    }
972
973    let selected = config
974        .get_option::<Vec<usize>>("selected_gpu_devices")
975        .unwrap_or_default();
976    let plan_raw = config.get_string_option("selected_layer_split_plan");
977    let parsed_plan =
978        if let Some(stages) = config.component_options.get("selected_layer_split_stages") {
979            crate::layer_split::parse_layer_split_stage_documents(stages)?
980        } else {
981            let plan_raw = plan_raw.as_deref().ok_or_else(|| {
982                FerrumError::config(
983                    "selected_distributed_strategy=layer_split requires selected_layer_split_plan",
984                )
985            })?;
986            crate::layer_split::parse_layer_split_plan(plan_raw)?
987        };
988    crate::layer_split::validate_layer_split_plan_for_devices(&parsed_plan, &selected)?;
989    if parsed_plan.total_layers() != num_layers {
990        return Err(FerrumError::config(format!(
991            "selected_layer_split_plan covers {} layers but model has {num_layers}",
992            parsed_plan.total_layers()
993        )));
994    }
995    Ok(Some(parsed_plan))
996}
997
998fn resolve_llama_layer_split_pipeline_mode(
999    config: &ComponentConfig,
1000    stage_count: usize,
1001) -> Result<ferrum_models::models::LlamaPipelineMode> {
1002    let Some(mode) = config.get_string_option("layer_split_pipeline_mode") else {
1003        return Ok(ferrum_models::models::LlamaPipelineMode::default_for_stage_count(stage_count));
1004    };
1005    let mode = ferrum_models::models::LlamaPipelineMode::from_config_value(&mode)?;
1006    if mode == ferrum_models::models::LlamaPipelineMode::Overlapped && stage_count != 2 {
1007        return Err(FerrumError::config(
1008            "layer_split_pipeline_mode=overlapped requires exactly two pipeline stages",
1009        ));
1010    }
1011    Ok(mode)
1012}
1013
1014#[cfg(test)]
1015fn resolve_llama_layer_stage_config(
1016    config: &ComponentConfig,
1017    num_layers: usize,
1018) -> Result<Option<ferrum_models::models::llama_family::LlamaFamilyLayerStageConfig>> {
1019    let Some(parsed_plan) = resolve_llama_layer_split_plan(config, num_layers)? else {
1020        return Ok(None);
1021    };
1022    let device_id = match &config.device {
1023        Device::CUDA(device_id) => *device_id,
1024        other => {
1025            return Err(FerrumError::unsupported(format!(
1026                "selected_distributed_strategy=layer_split requires a CUDA stage device, got {other:?}",
1027            )));
1028        }
1029    };
1030    parsed_plan
1031        .llama_stage_config_for_device(device_id)
1032        .map(Some)
1033}
1034
1035/// Generic LLM construction helper. Picks the model type (LlamaFamilyModel
1036/// or Qwen3MoeModel) by `arch`, opens a `NativeSafetensorsLoader<B>`, and
1037/// returns a `Box<dyn DecoderOnlyLLM>` ready to be wrapped in `LlmExecutor`.
1038///
1039/// Generic over both Dim 4 (`B`: hardware backend) and Dim 5 (`K`: KV
1040/// element type). Adding INT8 KV in PR C means adding the
1041/// `(Device::CUDA, KvCacheDtype::Int8) => build_llm::<CudaBackend, KvInt8>(...)`
1042/// arm to the cascade — this helper already accepts `K` so the callers
1043/// don't need to be touched again.
1044fn build_llm<B, K>(
1045    arch: ferrum_models::Architecture,
1046    qcfg: ferrum_models::models::LlamaFamilyConfig,
1047    moe_cfg: Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1048    model_path: &str,
1049    llama_layer_split_plan: Option<crate::layer_split::ParsedLayerSplitPlan>,
1050    llama_layer_split_pipeline_mode: Option<ferrum_models::models::LlamaPipelineMode>,
1051) -> Result<Box<dyn ferrum_models::common::DecoderOnlyLLM>>
1052where
1053    B: ferrum_kernels::backend::MoeLlmBackend,
1054    K: ferrum_kernels::backend::KvLayer<B>,
1055    ferrum_models::models::LlamaFamilyModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1056    ferrum_models::models::LlamaFamilyPipelineModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1057{
1058    if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1059        if llama_layer_split_plan.is_some() {
1060            return Err(FerrumError::unsupported(
1061                "CUDA layer_split stage loading is wired only for Llama-family dense models; \
1062                 Qwen3MoeModel requires a separate MoE stage loader.",
1063            ));
1064        }
1065        let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1066        let mc = moe_cfg.ok_or_else(|| {
1067            FerrumError::internal(
1068                "Qwen3Moe arch reached build_llm without Qwen3MoeConfig (caller bug)",
1069            )
1070        })?;
1071        Ok(Box::new(
1072            ferrum_models::models::Qwen3MoeModel::<B, K>::new_safetensors(mc, &weight_loader)?,
1073        ))
1074    } else {
1075        if llama_layer_split_plan.is_some() && !B::supports_device_ordinal_scope() {
1076            return Err(FerrumError::unsupported(
1077                "selected_distributed_strategy=layer_split requires a backend with \
1078                 device-scoped execution; refusing to silently use the default device",
1079            ));
1080        }
1081        let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1082        if let Some(plan) = llama_layer_split_plan {
1083            let stage_configs = plan.to_llama_stage_configs();
1084            let stage_device_ordinals = plan
1085                .stages
1086                .iter()
1087                .map(|stage| Some(stage.device))
1088                .collect::<Vec<_>>();
1089            let mut stages = Vec::with_capacity(stage_configs.len());
1090            for (idx, (stage_config, device_ordinal)) in stage_configs
1091                .into_iter()
1092                .zip(stage_device_ordinals.iter().copied())
1093                .enumerate()
1094            {
1095                tracing::info!(
1096                    "Loading Llama layer_split stage {idx} on backend device {:?}",
1097                    device_ordinal
1098                );
1099                let stage = B::with_device_ordinal(device_ordinal, || {
1100                    ferrum_models::models::LlamaFamilyModel::<B, K>::new_layer_stage(
1101                        qcfg.clone(),
1102                        &weight_loader,
1103                        stage_config,
1104                    )
1105                })?;
1106                stages.push(stage);
1107            }
1108            Ok(Box::new(ferrum_models::models::LlamaFamilyPipelineModel::<
1109                B,
1110                K,
1111            >::new_with_placement(
1112                stages,
1113                ferrum_models::models::LlamaPipelinePlacement::from_backend_device_ordinals(
1114                    stage_device_ordinals,
1115                )
1116                .with_pipeline_mode(
1117                    llama_layer_split_pipeline_mode.expect("layer split pipeline mode resolved"),
1118                ),
1119            )?))
1120        } else {
1121            Ok(Box::new(
1122                ferrum_models::models::LlamaFamilyModel::<B, K>::new(qcfg, &weight_loader)?,
1123            ))
1124        }
1125    }
1126}
1127
1128fn validate_registered_vnext_backend(
1129    kind: ferrum_models::vnext::ProductionExecutionKind,
1130    device: &Device,
1131    external_metadata_id: &ferrum_interfaces::vnext::ExternalModelMetadataId,
1132) -> Result<()> {
1133    use ferrum_models::vnext::ProductionExecutionKind;
1134
1135    match (kind, device) {
1136        (ProductionExecutionKind::CausalLanguage, Device::CUDA(_)) => {
1137            #[cfg(feature = "cuda")]
1138            return Ok(());
1139            #[cfg(not(feature = "cuda"))]
1140            Err(FerrumError::device(
1141                "registered vNext CUDA composition requires the 'cuda' feature",
1142            ))
1143        }
1144        #[cfg(any(target_os = "macos", target_os = "ios"))]
1145        (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1146            #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1147            return Ok(());
1148            #[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
1149            Err(FerrumError::device(
1150                "registered vNext Metal composition requires the 'metal' feature on an Apple platform",
1151            ))
1152        }
1153        (kind, device) => Err(FerrumError::unsupported(format!(
1154            "registered vNext model metadata {external_metadata_id} requires a {kind:?} backend composition, but {device} is not registered"
1155        ))),
1156    }
1157}
1158
1159fn create_registered_vnext_executor(
1160    config: &ComponentConfig,
1161    model_path: &std::path::Path,
1162    sources: Option<Arc<ProductionModelSourceBundle>>,
1163    prepared_model: Option<Arc<PreparedProductionModel>>,
1164    registration: ferrum_models::vnext::RegisteredProductionModel,
1165) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1166    use ferrum_models::vnext::ProductionExecutionKind;
1167
1168    validate_registered_vnext_backend(
1169        registration.execution_kind(),
1170        &config.device,
1171        registration.external_metadata_id(),
1172    )?;
1173
1174    let _prepared_model_reused = prepared_model.is_some();
1175    if let (Some(prepared), Some(sources)) = (prepared_model.as_ref(), sources.as_ref()) {
1176        if !Arc::ptr_eq(prepared.sources(), sources) {
1177            return Err(FerrumError::model(
1178                "prepared product model and component sources do not share one source lease",
1179            ));
1180        }
1181    }
1182    let prepared = match prepared_model {
1183        Some(prepared) => {
1184            if prepared.family().external_metadata_id() != registration.external_metadata_id() {
1185                return Err(FerrumError::model(format!(
1186                    "prepared product metadata {} differs from registered metadata {}",
1187                    prepared.family().external_metadata_id(),
1188                    registration.external_metadata_id()
1189                )));
1190            }
1191            prepared
1192        }
1193        None => Arc::new(match sources {
1194            Some(sources) => registration.prepare_from_sources(sources)?,
1195            None => registration.prepare(model_path)?,
1196        }),
1197    };
1198
1199    match (registration.execution_kind(), &config.device) {
1200        (ProductionExecutionKind::CausalLanguage, Device::CUDA(ordinal)) => {
1201            #[cfg(feature = "cuda")]
1202            {
1203                let model_info = prepared.model_info(
1204                    config.engine_config.model.model_id.clone(),
1205                    config.device.clone(),
1206                );
1207                let family = prepared.family();
1208                let family_fingerprint = family
1209                    .fingerprint()
1210                    .map_err(|error| FerrumError::model(error.to_string()))?;
1211                let program_fingerprint = family
1212                    .program()
1213                    .fingerprint()
1214                    .map_err(|error| FerrumError::model(error.to_string()))?;
1215                info!(
1216                    external_metadata_id = %registration.external_metadata_id(),
1217                    family_id = %family.family_id(),
1218                    family_fingerprint,
1219                    program_fingerprint,
1220                    prepared_model_reused = _prepared_model_reused,
1221                    backend = "cuda",
1222                    "Building registered model from a typed vNext execution plan"
1223                );
1224                let device_id = ferrum_interfaces::vnext::DeviceId::new(format!(
1225                    "device.cuda.{ordinal}"
1226                ))
1227                .map_err(|error| FerrumError::device(error.to_string()))?;
1228                let composition =
1229                    ferrum_kernels::backend::cuda::vnext_ops::CudaVNextComposition::create(
1230                        *ordinal,
1231                        device_id,
1232                        config.engine_config.runtime.attention_execution_policy,
1233                    )
1234                    .map_err(|error| {
1235                        FerrumError::device(format!("create vNext CUDA runtime: {error}"))
1236                    })?;
1237                let (
1238                    runtime,
1239                    operation_registry,
1240                    weight_materializers,
1241                    weight_materializer_id,
1242                    catalog,
1243                ) = composition.into_parts();
1244                let executor = crate::product_composition::create_vnext_executor(
1245                    &config.engine_config,
1246                    prepared.as_ref(),
1247                    model_info,
1248                    runtime,
1249                    operation_registry,
1250                    weight_materializers,
1251                    weight_materializer_id,
1252                    catalog,
1253                )?;
1254                info!(
1255                    resolved_plan_fingerprint = executor
1256                        .resolved_model_plan()
1257                        .map(|plan| plan.fingerprint())
1258                        .unwrap_or("missing"),
1259                    "Resolved product model plan is authoritative for vNext execution"
1260                );
1261                Ok(Arc::new(executor))
1262            }
1263            #[cfg(not(feature = "cuda"))]
1264            {
1265                let _ = (ordinal, model_path, prepared);
1266                Err(FerrumError::device(
1267                    "registered vNext CUDA composition requires the 'cuda' feature",
1268                ))
1269            }
1270        }
1271        #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1272        (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1273            let model_info = prepared.model_info(
1274                config.engine_config.model.model_id.clone(),
1275                config.device.clone(),
1276            );
1277            let family = prepared.family();
1278            let family_fingerprint = family
1279                .fingerprint()
1280                .map_err(|error| FerrumError::model(error.to_string()))?;
1281            let program_fingerprint = family
1282                .program()
1283                .fingerprint()
1284                .map_err(|error| FerrumError::model(error.to_string()))?;
1285            info!(
1286                external_metadata_id = %registration.external_metadata_id(),
1287                family_id = %family.family_id(),
1288                family_fingerprint,
1289                program_fingerprint,
1290                prepared_model_reused = _prepared_model_reused,
1291                backend = "metal",
1292                "Building registered model from a typed vNext execution plan"
1293            );
1294            let device_id = ferrum_interfaces::vnext::DeviceId::new("device.metal.0")
1295                .map_err(|error| FerrumError::device(error.to_string()))?;
1296            let composition =
1297                ferrum_kernels::backend::metal::vnext_ops::MetalVNextComposition::create(
1298                    device_id,
1299                )
1300                .map_err(|error| {
1301                    FerrumError::device(format!("create vNext Metal runtime: {error}"))
1302                })?;
1303            let (
1304                runtime,
1305                operation_registry,
1306                weight_materializers,
1307                weight_materializer_id,
1308                catalog,
1309            ) = composition.into_parts();
1310            let executor = crate::product_composition::create_vnext_executor(
1311                &config.engine_config,
1312                prepared.as_ref(),
1313                model_info,
1314                runtime,
1315                operation_registry,
1316                weight_materializers,
1317                weight_materializer_id,
1318                catalog,
1319            )?;
1320            info!(
1321                resolved_plan_fingerprint = executor
1322                    .resolved_model_plan()
1323                    .map(|plan| plan.fingerprint())
1324                    .unwrap_or("missing"),
1325                "Resolved product model plan is authoritative for vNext execution"
1326            );
1327            Ok(Arc::new(executor))
1328        }
1329        (kind, device) => Err(FerrumError::unsupported(format!(
1330            "registered vNext model metadata {} requires a {kind:?} backend composition, but {device} is not registered",
1331            registration.external_metadata_id()
1332        ))),
1333    }
1334}
1335
1336/// Back-compat alias; retained so external test fixtures referencing
1337/// the old name continue to compile while we sweep call sites.
1338#[deprecated(note = "use `LlmExecutorFactory` (renamed PR A — Dim 1/3 cleanup)")]
1339pub type CandleExecutorFactory = LlmExecutorFactory;
1340
1341#[async_trait]
1342impl ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>> for LlmExecutorFactory {
1343    async fn create(
1344        &self,
1345        config: &ComponentConfig,
1346    ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1347        use candle_core::{DType, Device as CandleDevice};
1348        use ferrum_models::weight_format::WeightFormat;
1349
1350        // Try to load model from path
1351        let checkpoint_capture_enabled = config
1352            .engine_config
1353            .runtime
1354            .vnext_checkpoint_capture
1355            .is_some();
1356        let model_sources = config.model_sources.clone();
1357        let prepared_model = config.prepared_model.clone();
1358        let model_path = model_sources
1359            .as_ref()
1360            .map(|sources| sources.weights().path().display().to_string())
1361            .or_else(|| config.get_string_option("model_path"))
1362            .or_else(|| config.engine_config.runtime.model_path.clone());
1363
1364        let model_path = match model_path {
1365            Some(path) => path,
1366            None => {
1367                if checkpoint_capture_enabled {
1368                    return Err(FerrumError::unsupported(
1369                        "vNext checkpoint capture requires a registered model source",
1370                    ));
1371                }
1372                info!("No model path found, falling back to stub executor");
1373                return StubExecutorFactory.create(config).await;
1374            }
1375        };
1376
1377        // Dim 3 dispatch — peer-level enum (no GGUF special-case at the
1378        // top of an architecture cascade). Adding a new format (AWQ /
1379        // EXL2 / HQQ) means a new `WeightFormat` variant + a matching
1380        // `WeightLoader<B>` impl in `ferrum-quantization`.
1381        let weight_fmt = WeightFormat::detect(std::path::Path::new(&model_path))?;
1382        info!(
1383            "Loading model from {} (format: {})",
1384            model_path,
1385            weight_fmt.label()
1386        );
1387
1388        // Registered vNext packages resolve from immutable external metadata
1389        // before every legacy weight-format loader. A typed GGUF source carries
1390        // its semantic metadata separately, so routing it through the generic
1391        // GGUF loader first would silently bypass the registered vNext package.
1392        // Direct GGUF paths with colocated semantic/tokenizer metadata can
1393        // carry the same typed bundle. Migrated architectures without that
1394        // metadata fail closed during product source resolution; only explicit
1395        // legacy registry rows may continue through the generic GGUF loader.
1396        let production_registration = match model_sources.as_ref() {
1397            Some(sources) => Some(ferrum_models::vnext::resolve_registered_model_from_sources(
1398                sources,
1399            )?),
1400            None if !matches!(weight_fmt, WeightFormat::Gguf { .. }) => {
1401                Some(ferrum_models::vnext::resolve_registered_model_from_dir(
1402                    std::path::Path::new(&model_path),
1403                )?)
1404            }
1405            None => None,
1406        };
1407        if let Some(production_registration) = production_registration {
1408            match production_registration {
1409                ferrum_models::vnext::ProductionModelRegistration::Registered(registration) => {
1410                    return create_registered_vnext_executor(
1411                        config,
1412                        std::path::Path::new(&model_path),
1413                        model_sources,
1414                        prepared_model,
1415                        registration,
1416                    );
1417                }
1418                ferrum_models::vnext::ProductionModelRegistration::LegacyRegistered {
1419                    external_metadata_id,
1420                } => {
1421                    if checkpoint_capture_enabled {
1422                        return Err(FerrumError::unsupported(format!(
1423                            "vNext checkpoint capture requires a migrated model; metadata {external_metadata_id} is still legacy"
1424                        )));
1425                    }
1426                    info!(
1427                        %external_metadata_id,
1428                        "Entering the explicitly registered legacy model path"
1429                    );
1430                }
1431            }
1432        }
1433        if checkpoint_capture_enabled {
1434            return Err(FerrumError::unsupported(
1435                "vNext checkpoint capture requires a registered vNext model package",
1436            ));
1437        }
1438
1439        if let WeightFormat::Gguf { ref path } = weight_fmt {
1440            // Legacy or direct-path GGUF packages still use the monolithic
1441            // loader. Registered typed packages have already returned above.
1442            let (llm, model_info) = ferrum_models::gguf_engine_loader::load_gguf_decoder_with_info(
1443                path,
1444                &config.device,
1445                config.engine_config.model.model_id.clone(),
1446            )?;
1447            return Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)));
1448        }
1449
1450        // Explicitly registered legacy safetensors and the diagnostic CPU
1451        // reference path continue through the old registry until each family
1452        // migration deletes its legacy row and architecture branch together.
1453        let mut config_manager = ferrum_models::ConfigManager::new();
1454        let model_def = config_manager
1455            .load_from_path(std::path::Path::new(&model_path))
1456            .await?;
1457
1458        info!(
1459            "  Architecture: {:?}, Layers: {}, Vocab: {}",
1460            model_def.architecture, model_def.num_hidden_layers, model_def.vocab_size
1461        );
1462
1463        // Dense decoder families below use Ferrum's Backend<B> stack and do
1464        // not need a Candle device. Resolve the legacy Candle device lazily so
1465        // the official CUDA graph can run those families without enabling the
1466        // explicit candle-cuda-compat feature. Candle-backed BERT/CLIP/Whisper
1467        // arms still fail closed on CUDA unless that compatibility feature is
1468        // intentionally selected.
1469        let legacy_candle_device = || -> Result<CandleDevice> {
1470            match &config.device {
1471                Device::CPU => Ok(CandleDevice::Cpu),
1472                #[cfg(feature = "candle-cuda-compat")]
1473                Device::CUDA(id) => CandleDevice::new_cuda(*id)
1474                    .map_err(|e| FerrumError::device(format!("CUDA error: {}", e))),
1475                #[cfg(not(feature = "candle-cuda-compat"))]
1476                Device::CUDA(_) => Err(FerrumError::unsupported(
1477                    "legacy Candle CUDA executors require the candle-cuda-compat feature",
1478                )),
1479                #[cfg(any(target_os = "macos", target_os = "ios"))]
1480                Device::Metal => CandleDevice::new_metal(0)
1481                    .map_err(|e| FerrumError::device(format!("Metal error: {}", e))),
1482                Device::ROCm(_) => Err(FerrumError::device("ROCm not yet supported")),
1483            }
1484        };
1485
1486        // Select dtype.
1487        //
1488        // IMPORTANT (correctness): Metal defaults to FP32 for stability.
1489        // This was introduced to address cases where CPU inference is correct but Metal results
1490        // can deviate. If you want higher performance and accept potential numerical risk, you
1491        // can explicitly opt-in via env:
1492        // - FERRUM_METAL_DTYPE=fp16|fp32 (takes precedence on Metal)
1493        // - FERRUM_DTYPE=fp16|fp32 (global override for non-CPU)
1494        let dtype: DType = RegistryRuntimeEnv::from_runtime_knobs(&config.engine_config.runtime)
1495            .dtype_for_device(&config.device);
1496
1497        // Create model based on architecture
1498        info!("Building model...");
1499        match model_def.architecture {
1500            // All Llama-family decoders (Llama / Llama-2 / Llama-3 / Qwen2 /
1501            // Qwen2.5 / Qwen3) share `LlamaFamilyModel<B>` + `LlmExecutor`.
1502            // Only the config constructor differs.
1503            arch @ (ferrum_models::Architecture::Llama
1504            | ferrum_models::Architecture::Qwen2
1505            | ferrum_models::Architecture::Qwen3
1506            | ferrum_models::Architecture::Qwen3Moe
1507            | ferrum_models::Architecture::Gemma3
1508            | ferrum_models::Architecture::Mistral) => {
1509                let _loader = ferrum_models::SafeTensorsLoader::new(&model_path);
1510                let model_dir_path: std::path::PathBuf = model_path.clone().into();
1511
1512                // Tensor parallelism (FERRUM_TP>1) was wired against the
1513                // pre-Architecture-v2 CandleBackend and hasn't been ported
1514                // to the Backend<B> trait stack. Reject explicitly so users
1515                // don't get a silent single-GPU fallback.
1516                let tp_size = config.engine_config.runtime.tp.unwrap_or(0);
1517                if tp_size > 1 {
1518                    return Err(FerrumError::unsupported(
1519                        "FERRUM_TP>1 not supported on the Backend<B> path. \
1520                         Run with FERRUM_TP=1 (default) for single-GPU inference.",
1521                    ));
1522                }
1523                // GPTQ is loaded via NativeSafetensorsLoader's load_linear:
1524                // it auto-detects `<name>.qweight` tensors and constructs
1525                // a GptqLinear via Backend::load_gptq. The QuantizeConfig
1526                // probe is kept for back-compat with old configs that
1527                // explicitly listed a quantization method.
1528                let _ = ferrum_models::loader::QuantizeConfig::from_model_dir(&model_dir_path);
1529
1530                // Resolve the architecture-specific config in one place.
1531                // Qwen3-MoE needs both `LlamaFamilyConfig` (dense attention
1532                // shares the Llama path) AND `Qwen3MoeConfig` (router +
1533                // experts); other arches only need `LlamaFamilyConfig`.
1534                let (qcfg, moe_cfg): (
1535                    ferrum_models::models::LlamaFamilyConfig,
1536                    Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1537                ) = match arch {
1538                    ferrum_models::Architecture::Qwen3Moe => {
1539                        info!("Loading Qwen3-MoE via Qwen3MoeModel (safetensors GPTQ)");
1540                        let mc = ferrum_models::moe_config::Qwen3MoeConfig::from_def(&model_def)?;
1541                        // dense attention reuses the Qwen3 dense config —
1542                        // build_llm only consumes `qcfg` on the LlamaFamily
1543                        // arm, so passing the MoE's `base` is fine.
1544                        (mc.base.clone(), Some(mc))
1545                    }
1546                    ferrum_models::Architecture::Qwen3 => {
1547                        info!("Loading Qwen3 via LlamaFamilyModel (QK-norm on)");
1548                        (
1549                            ferrum_models::models::LlamaFamilyConfig::qwen3_from_def(&model_def),
1550                            None,
1551                        )
1552                    }
1553                    ferrum_models::Architecture::Qwen2 => {
1554                        info!("Loading Qwen2 via LlamaFamilyModel");
1555                        (
1556                            ferrum_models::models::LlamaFamilyConfig::qwen2_from_def(&model_def),
1557                            None,
1558                        )
1559                    }
1560                    ferrum_models::Architecture::Mistral => {
1561                        info!("Loading Mistral via LlamaFamilyModel (sliding_window from config)");
1562                        (
1563                            ferrum_models::models::LlamaFamilyConfig::mistral_from_def(&model_def),
1564                            None,
1565                        )
1566                    }
1567                    ferrum_models::Architecture::Gemma3 => {
1568                        info!(
1569                            "Loading Gemma3 via LlamaFamilyModel (5:1 SWA, dual rope, GeGLU, \
1570                             sandwich norms)"
1571                        );
1572                        (
1573                            ferrum_models::models::LlamaFamilyConfig::gemma3_from_def(&model_def),
1574                            None,
1575                        )
1576                    }
1577                    _ => {
1578                        info!("Loading Llama via LlamaFamilyModel");
1579                        (
1580                            ferrum_models::models::LlamaFamilyConfig::llama_from_def(&model_def),
1581                            None,
1582                        )
1583                    }
1584                };
1585
1586                let model_info =
1587                    model_def.to_model_info(config.engine_config.model.model_id.to_string());
1588                let llama_layer_split_plan =
1589                    resolve_llama_layer_split_plan(config, qcfg.num_layers)?;
1590                let llama_layer_split_pipeline_mode = llama_layer_split_plan
1591                    .as_ref()
1592                    .map(|plan| resolve_llama_layer_split_pipeline_mode(config, plan.stages.len()))
1593                    .transpose()?;
1594
1595                // (Dim 4, Dim 5) cascade: pick `B` from device, `K` from
1596                // kv-dtype, and dispatch to the generic `build_llm` helper.
1597                // PR C will extend this with `(CUDA, Int8) => build_llm::<CudaBackend, KvInt8>(...)`
1598                // — model wire-up already accepts `K`, so adding INT8 only
1599                // touches this match.
1600                use ferrum_interfaces::kv_dtype::KvFp16;
1601                #[cfg(feature = "cuda")]
1602                use ferrum_interfaces::kv_dtype::KvInt8;
1603                use ferrum_types::KvCacheDtype;
1604                let kv_dtype = config.engine_config.kv_cache.dtype;
1605                let llm: Box<dyn ferrum_models::common::DecoderOnlyLLM> =
1606                    match (&config.device, kv_dtype) {
1607                        (Device::CPU, KvCacheDtype::Fp16) => {
1608                            info!("  Backend: CPU, KV: fp16");
1609                            build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
1610                                arch,
1611                                qcfg,
1612                                moe_cfg,
1613                                &model_path,
1614                                llama_layer_split_plan,
1615                                llama_layer_split_pipeline_mode,
1616                            )?
1617                        }
1618                        #[cfg(any(target_os = "macos", target_os = "ios"))]
1619                        (Device::Metal, KvCacheDtype::Fp16) => {
1620                            #[cfg(feature = "metal")]
1621                            {
1622                                // FERRUM_METAL_DTYPE=f16 toggles fp16 weight storage
1623                                // inside MetalBackend. Halves big-tensor RAM;
1624                                // recommended for 4B+ models on 16 GB Macs.
1625                                let dtype_hint = RegistryRuntimeEnv::from_runtime_knobs(
1626                                    &config.engine_config.runtime,
1627                                )
1628                                .metal_dtype_hint();
1629                                info!("  Backend: Metal (weights {}), KV: fp16", dtype_hint);
1630                                build_llm::<ferrum_kernels::backend::metal::MetalBackend, KvFp16>(
1631                                    arch,
1632                                    qcfg,
1633                                    moe_cfg,
1634                                    &model_path,
1635                                    llama_layer_split_plan,
1636                                    llama_layer_split_pipeline_mode,
1637                                )?
1638                            }
1639                            #[cfg(not(feature = "metal"))]
1640                            {
1641                                return Err(FerrumError::device(
1642                                    "Metal requested but 'metal' feature not enabled",
1643                                ));
1644                            }
1645                        }
1646                        (Device::CUDA(_), KvCacheDtype::Fp16) => {
1647                            #[cfg(feature = "cuda")]
1648                            {
1649                                info!("  Backend: CUDA, KV: fp16");
1650                                build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvFp16>(
1651                                    arch,
1652                                    qcfg,
1653                                    moe_cfg,
1654                                    &model_path,
1655                                    llama_layer_split_plan,
1656                                    llama_layer_split_pipeline_mode,
1657                                )?
1658                            }
1659                            #[cfg(not(feature = "cuda"))]
1660                            {
1661                                return Err(FerrumError::device(
1662                                    "CUDA requested but 'cuda' feature not enabled",
1663                                ));
1664                            }
1665                        }
1666                        (Device::CUDA(_), KvCacheDtype::Int8) => {
1667                            #[cfg(feature = "cuda")]
1668                            {
1669                                // Dim 5 PR C: only CudaBackend implements
1670                                // BackendInt8KvOps with real launchers; LlamaFamilyModel<B, KvInt8>
1671                                // uses LayerKvCache::Int8 + paged INT8 KV. Qwen3-MoE
1672                                // INT8 KV is a follow-up — reject here so users
1673                                // see a clear error instead of a panic at
1674                                // alloc_paged_int8_layer time.
1675                                if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1676                                    return Err(FerrumError::unsupported(
1677                                        "INT8 KV cache is not yet wired through Qwen3MoeModel \
1678                                     (LlamaFamilyModel-only in PR C). Use --kv-dtype fp16 \
1679                                     for MoE models or wait for the follow-up PR.",
1680                                    ));
1681                                }
1682                                info!("  Backend: CUDA, KV: int8 (paged, vLLM-style)");
1683                                build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvInt8>(
1684                                    arch,
1685                                    qcfg,
1686                                    moe_cfg,
1687                                    &model_path,
1688                                    llama_layer_split_plan,
1689                                    llama_layer_split_pipeline_mode,
1690                                )?
1691                            }
1692                            #[cfg(not(feature = "cuda"))]
1693                            {
1694                                return Err(FerrumError::device(
1695                                    "CUDA requested but 'cuda' feature not enabled",
1696                                ));
1697                            }
1698                        }
1699                        (dev, dt) => {
1700                            return Err(FerrumError::unsupported(format!(
1701                                "(device={dev:?}, kv_dtype={dt:?}) not implemented — \
1702                             see docs/dim5-model-wireup-plan.md"
1703                            )));
1704                        }
1705                    };
1706
1707                Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)))
1708            }
1709            ferrum_models::Architecture::Bert => {
1710                info!("Using BERT executor for embeddings");
1711                let executor = ferrum_models::BertModelExecutor::from_path(
1712                    &model_path,
1713                    &model_def,
1714                    legacy_candle_device()?,
1715                )
1716                .await?;
1717
1718                Ok(Arc::new(executor))
1719            }
1720            ferrum_models::Architecture::Clip => {
1721                info!("Using CLIP executor for multimodal embeddings");
1722                let executor = ferrum_models::ClipModelExecutor::from_path(
1723                    &model_path,
1724                    legacy_candle_device()?,
1725                    dtype,
1726                )?;
1727                Ok(Arc::new(executor))
1728            }
1729            ferrum_models::Architecture::Whisper => {
1730                info!("Using Whisper executor for ASR");
1731                let executor = ferrum_models::WhisperModelExecutor::from_path(
1732                    &model_path,
1733                    legacy_candle_device()?,
1734                    dtype,
1735                )?;
1736                Ok(Arc::new(executor))
1737            }
1738            _ => Err(FerrumError::model(format!(
1739                "Architecture {:?} not supported",
1740                model_def.architecture
1741            ))),
1742        }
1743    }
1744
1745    fn metadata(&self) -> ComponentMetadata {
1746        ComponentMetadata {
1747            name: "llm".to_string(),
1748            version: "0.2.0".to_string(),
1749            description: "LLM executor (LlamaFamily / Qwen3MoE via Backend<B>; \
1750                 BERT / CLIP / Whisper via candle)"
1751                .to_string(),
1752            supported_devices: cpu_cuda_and_optional_metal_devices(),
1753            capabilities: vec![
1754                "llama".to_string(),
1755                "qwen2".to_string(),
1756                "qwen3".to_string(),
1757                "qwen3_moe".to_string(),
1758                "mistral".to_string(),
1759                "bert".to_string(),
1760                "clip".to_string(),
1761                "whisper".to_string(),
1762                "safetensors".to_string(),
1763                "gguf".to_string(),
1764                "fp16".to_string(),
1765                "fp32".to_string(),
1766            ],
1767        }
1768    }
1769}
1770
1771// ============================================================================
1772// Stub Tokenizer Implementation
1773// ============================================================================
1774
1775/// Stub tokenizer for testing
1776pub struct StubTokenizer {
1777    vocab_size: usize,
1778    info: ferrum_interfaces::TokenizerInfo,
1779}
1780
1781impl StubTokenizer {
1782    /// Create a new stub tokenizer
1783    pub fn new(vocab_size: usize) -> Self {
1784        let info = ferrum_interfaces::TokenizerInfo {
1785            tokenizer_type: ferrum_interfaces::tokenizer::TokenizerType::BPE,
1786            vocab_size,
1787            special_tokens: ferrum_types::SpecialTokens::default(),
1788            supports_incremental: false,
1789            supports_chat_template: false,
1790            max_token_length: None,
1791            model_name: Some("stub".into()),
1792        };
1793
1794        Self { vocab_size, info }
1795    }
1796}
1797
1798impl std::fmt::Debug for StubTokenizer {
1799    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1800        f.debug_struct("StubTokenizer")
1801            .field("vocab_size", &self.vocab_size)
1802            .finish()
1803    }
1804}
1805
1806impl Tokenizer for StubTokenizer {
1807    fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<ferrum_types::TokenId>> {
1808        let tokens: Vec<ferrum_types::TokenId> = text
1809            .split_whitespace()
1810            .enumerate()
1811            .map(|(i, _)| ferrum_types::TokenId::new((i % self.vocab_size) as u32))
1812            .collect();
1813
1814        Ok(if tokens.is_empty() {
1815            vec![ferrum_types::TokenId::new(0)]
1816        } else {
1817            tokens
1818        })
1819    }
1820
1821    fn decode(&self, tokens: &[ferrum_types::TokenId], _skip_special: bool) -> Result<String> {
1822        Ok(tokens
1823            .iter()
1824            .map(|t| format!("token_{}", t.get()))
1825            .collect::<Vec<_>>()
1826            .join(" "))
1827    }
1828
1829    fn decode_incremental(
1830        &self,
1831        _prev: &[ferrum_types::TokenId],
1832        next: ferrum_types::TokenId,
1833    ) -> Result<String> {
1834        Ok(format!("token_{} ", next.get()))
1835    }
1836
1837    fn vocab_size(&self) -> usize {
1838        self.vocab_size
1839    }
1840
1841    fn special_tokens(&self) -> &ferrum_types::SpecialTokens {
1842        &self.info.special_tokens
1843    }
1844
1845    fn token_id(&self, _text: &str) -> Option<ferrum_types::TokenId> {
1846        Some(ferrum_types::TokenId::new(0))
1847    }
1848
1849    fn token_text(&self, _token_id: ferrum_types::TokenId) -> Option<&str> {
1850        None
1851    }
1852
1853    fn info(&self) -> ferrum_interfaces::TokenizerInfo {
1854        self.info.clone()
1855    }
1856}
1857
1858// ============================================================================
1859// Global Registry
1860// ============================================================================
1861
1862/// Global registry instance
1863static GLOBAL_REGISTRY: OnceLock<Arc<ComponentRegistry>> = OnceLock::new();
1864
1865/// Get the global registry, initializing with defaults if needed
1866pub fn global_registry() -> Arc<ComponentRegistry> {
1867    GLOBAL_REGISTRY
1868        .get_or_init(|| {
1869            info!("Initializing global component registry");
1870            Arc::new(ComponentRegistry::with_defaults())
1871        })
1872        .clone()
1873}
1874
1875/// Initialize global registry with a custom registry
1876pub fn set_global_registry(registry: Arc<ComponentRegistry>) -> Result<()> {
1877    GLOBAL_REGISTRY
1878        .set(registry)
1879        .map_err(|_| FerrumError::internal("Global registry already initialized"))
1880}
1881
1882// ============================================================================
1883// Tests
1884// ============================================================================
1885
1886#[cfg(test)]
1887mod tests {
1888    use super::*;
1889    use safetensors::tensor::{serialize_to_file, Dtype, TensorView};
1890    use std::path::{Path, PathBuf};
1891
1892    fn unique_test_dir(name: &str) -> PathBuf {
1893        let mut dir = std::env::temp_dir();
1894        dir.push(format!(
1895            "ferrum-{name}-{}-{}",
1896            std::process::id(),
1897            std::time::SystemTime::now()
1898                .duration_since(std::time::UNIX_EPOCH)
1899                .unwrap()
1900                .as_nanos()
1901        ));
1902        std::fs::create_dir_all(&dir).unwrap();
1903        dir
1904    }
1905
1906    fn write_qwen35_fixture_config(dir: &Path) {
1907        let config = serde_json::json!({
1908            "architectures": ["Qwen3_5ForConditionalGeneration"],
1909            "model_type": "qwen3_5",
1910            "vocab_size": 3,
1911            "max_position_embeddings": 16,
1912            "rms_norm_eps": 1e-6,
1913            "rope_theta": 10000.0,
1914            "tie_word_embeddings": false,
1915            "text_config": {
1916                "model_type": "qwen3_5_text",
1917                "hidden_size": 2,
1918                "intermediate_size": 2,
1919                "num_hidden_layers": 2,
1920                "layer_types": ["linear_attention", "full_attention"],
1921                "linear_num_key_heads": 1,
1922                "linear_num_value_heads": 1,
1923                "linear_key_head_dim": 1,
1924                "linear_value_head_dim": 1,
1925                "linear_conv_kernel_dim": 1,
1926                "mamba_ssm_dtype": "float32",
1927                "head_dim": 2,
1928                "num_attention_heads": 1,
1929                "num_key_value_heads": 1,
1930                "vocab_size": 3,
1931                "max_position_embeddings": 16,
1932                "tie_word_embeddings": false
1933            }
1934        });
1935        std::fs::write(
1936            dir.join("config.json"),
1937            serde_json::to_string_pretty(&config).unwrap(),
1938        )
1939        .unwrap();
1940    }
1941
1942    fn write_qwen35_fixture_safetensors(dir: &Path) {
1943        let tensors: Vec<(String, Vec<f32>)> = vec![
1944            (
1945                "model.embed_tokens.weight".to_string(),
1946                vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1947            ),
1948            ("model.norm.weight".to_string(), vec![0.0, 0.0]),
1949            (
1950                "model.lm_head.weight".to_string(),
1951                vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1952            ),
1953            (
1954                "model.layers.0.input_layernorm.weight".to_string(),
1955                vec![0.0, 0.0],
1956            ),
1957            (
1958                "model.layers.0.post_attention_layernorm.weight".to_string(),
1959                vec![0.0, 0.0],
1960            ),
1961            (
1962                "model.layers.0.linear_attn.in_proj_qkv.weight".to_string(),
1963                vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1964            ),
1965            (
1966                "model.layers.0.linear_attn.in_proj_z.weight".to_string(),
1967                vec![1.0, -1.0],
1968            ),
1969            (
1970                "model.layers.0.linear_attn.in_proj_b.weight".to_string(),
1971                vec![0.5, 0.25],
1972            ),
1973            (
1974                "model.layers.0.linear_attn.in_proj_a.weight".to_string(),
1975                vec![-0.25, 0.75],
1976            ),
1977            (
1978                "model.layers.0.linear_attn.conv1d.weight".to_string(),
1979                vec![1.0, 1.0, 1.0],
1980            ),
1981            ("model.layers.0.linear_attn.A_log".to_string(), vec![0.0]),
1982            ("model.layers.0.linear_attn.dt_bias".to_string(), vec![0.0]),
1983            (
1984                "model.layers.0.linear_attn.norm.weight".to_string(),
1985                vec![1.0],
1986            ),
1987            (
1988                "model.layers.0.linear_attn.out_proj.weight".to_string(),
1989                vec![1.0, -0.5],
1990            ),
1991            (
1992                "model.layers.0.mlp.gate_proj.weight".to_string(),
1993                vec![0.2, 0.1, -0.1, 0.3],
1994            ),
1995            (
1996                "model.layers.0.mlp.up_proj.weight".to_string(),
1997                vec![0.4, -0.2, 0.3, 0.5],
1998            ),
1999            (
2000                "model.layers.0.mlp.down_proj.weight".to_string(),
2001                vec![1.0, 0.0, 0.0, 1.0],
2002            ),
2003            (
2004                "model.layers.1.input_layernorm.weight".to_string(),
2005                vec![0.0, 0.0],
2006            ),
2007            (
2008                "model.layers.1.post_attention_layernorm.weight".to_string(),
2009                vec![0.0, 0.0],
2010            ),
2011            (
2012                "model.layers.1.self_attn.q_proj.weight".to_string(),
2013                vec![1.0, 0.0, 0.0, 1.0],
2014            ),
2015            (
2016                "model.layers.1.self_attn.k_proj.weight".to_string(),
2017                vec![0.5, 0.0, 0.0, 0.5],
2018            ),
2019            (
2020                "model.layers.1.self_attn.v_proj.weight".to_string(),
2021                vec![1.0, 1.0, -0.5, 0.5],
2022            ),
2023            (
2024                "model.layers.1.self_attn.o_proj.weight".to_string(),
2025                vec![1.0, 0.0, 0.0, 1.0],
2026            ),
2027            (
2028                "model.layers.1.self_attn.q_norm.weight".to_string(),
2029                vec![1.0, 1.0],
2030            ),
2031            (
2032                "model.layers.1.self_attn.k_norm.weight".to_string(),
2033                vec![1.0, 1.0],
2034            ),
2035            (
2036                "model.layers.1.mlp.gate_proj.weight".to_string(),
2037                vec![-0.2, 0.2, 0.1, 0.3],
2038            ),
2039            (
2040                "model.layers.1.mlp.up_proj.weight".to_string(),
2041                vec![0.25, 0.5, -0.3, 0.4],
2042            ),
2043            (
2044                "model.layers.1.mlp.down_proj.weight".to_string(),
2045                vec![0.5, 0.25, -0.2, 0.75],
2046            ),
2047        ];
2048        let views = tensors
2049            .into_iter()
2050            .map(|(name, values)| {
2051                let bytes = values
2052                    .iter()
2053                    .flat_map(|value| value.to_le_bytes())
2054                    .collect::<Vec<_>>()
2055                    .into_boxed_slice();
2056                let bytes: &'static [u8] = Box::leak(bytes);
2057                (
2058                    name,
2059                    TensorView::new(Dtype::F32, vec![values.len()], bytes).unwrap(),
2060                )
2061            })
2062            .collect::<Vec<_>>();
2063        serialize_to_file(
2064            views,
2065            &None::<std::collections::HashMap<String, String>>,
2066            &dir.join("model.safetensors"),
2067        )
2068        .unwrap();
2069    }
2070
2071    fn write_qwen35_fixture_model_dir() -> PathBuf {
2072        let dir = unique_test_dir("qwen35-vnext-route");
2073        write_qwen35_fixture_config(&dir);
2074        write_qwen35_fixture_safetensors(&dir);
2075        dir
2076    }
2077
2078    fn qwen35_fixture_component_config(model_dir: &Path) -> ComponentConfig {
2079        let mut engine_config = EngineConfig::default();
2080        engine_config.backend.device = Device::CPU;
2081        engine_config.backend.backend_options.insert(
2082            "model_path".to_string(),
2083            serde_json::Value::String(model_dir.to_string_lossy().to_string()),
2084        );
2085        ComponentConfig::from_engine_config(&engine_config)
2086    }
2087
2088    fn enable_checkpoint_capture(config: &mut ComponentConfig, output_dir: PathBuf) {
2089        config.engine_config.runtime.vnext_checkpoint_capture =
2090            Some(ferrum_types::VNextCheckpointCaptureConfig {
2091                output_dir,
2092                value_ids: vec!["value.output.logits".to_owned()],
2093                maximum_prefill_waves: 1,
2094                maximum_decode_waves: 0,
2095                capture_product_output: false,
2096                teacher_forcing: None,
2097            });
2098    }
2099
2100    #[test]
2101    fn checkpoint_capture_rejects_stub_fallback() {
2102        let mut config = ComponentConfig::from_engine_config(&EngineConfig::default());
2103        enable_checkpoint_capture(&mut config, PathBuf::from("capture"));
2104
2105        let error = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2106            Ok(_) => panic!("checkpoint capture unexpectedly entered the stub executor"),
2107            Err(error) => error.to_string(),
2108        };
2109        assert!(
2110            error.contains("requires a registered model source"),
2111            "{error}"
2112        );
2113    }
2114
2115    #[test]
2116    fn qwen35_registry_routes_registered_package_before_legacy_loading() {
2117        let dir = write_qwen35_fixture_model_dir();
2118        std::fs::remove_file(dir.join("model.safetensors")).unwrap();
2119        let config = qwen35_fixture_component_config(&dir);
2120
2121        let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2122            Ok(_) => panic!("registered Qwen3.5 package unexpectedly accepted CPU execution"),
2123            Err(err) => err.to_string(),
2124        };
2125
2126        assert!(
2127            err.contains(ferrum_models::vnext::qwen35::EXTERNAL_METADATA_ID),
2128            "{err}"
2129        );
2130        assert!(err.contains("but cpu is not registered"), "{err}");
2131        let _ = std::fs::remove_dir_all(dir);
2132    }
2133
2134    #[test]
2135    fn qwen35_typed_gguf_routes_registered_package_before_legacy_gguf_loading() {
2136        let dir = unique_test_dir("qwen35-typed-gguf-route");
2137        write_qwen35_fixture_config(&dir);
2138        std::fs::write(dir.join("tokenizer.json"), br#"{"version":"1.0"}"#).unwrap();
2139        std::fs::write(
2140            dir.join("tokenizer_config.json"),
2141            br#"{"chat_template":"fixture"}"#,
2142        )
2143        .unwrap();
2144        let gguf = dir.join("model.gguf");
2145        std::fs::write(&gguf, b"not-a-real-gguf").unwrap();
2146        let original = ferrum_interfaces::vnext::OriginalModelSource {
2147            kind: ferrum_interfaces::vnext::ModelSourceKind::LocalDirectory,
2148            location: dir.display().to_string(),
2149            requested_revision: None,
2150        };
2151        let sources = Arc::new(
2152            ProductionModelSourceBundle::open(
2153                &dir,
2154                &dir,
2155                ferrum_models::vnext::ProductionWeightArtifact::gguf_file(&gguf),
2156                ferrum_interfaces::vnext::OriginalModelSources {
2157                    semantic: original.clone(),
2158                    tokenizer: original.clone(),
2159                    weights: ferrum_interfaces::vnext::OriginalModelSource {
2160                        kind: ferrum_interfaces::vnext::ModelSourceKind::LocalFile,
2161                        location: gguf.display().to_string(),
2162                        requested_revision: None,
2163                    },
2164                },
2165            )
2166            .unwrap(),
2167        );
2168        let mut engine = EngineConfig::default();
2169        engine.backend.device = Device::CPU;
2170        let mut config = ComponentConfig::from_engine_config(&engine);
2171        config.model_sources = Some(sources);
2172
2173        let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2174            Ok(_) => panic!("registered Qwen3.5 GGUF unexpectedly accepted CPU execution"),
2175            Err(err) => err.to_string(),
2176        };
2177
2178        assert!(
2179            err.contains(ferrum_models::vnext::qwen35::EXTERNAL_METADATA_ID),
2180            "{err}"
2181        );
2182        assert!(err.contains("but cpu is not registered"), "{err}");
2183        assert!(!err.contains("GGUF"), "{err}");
2184        let _ = std::fs::remove_dir_all(dir);
2185    }
2186
2187    #[test]
2188    fn test_registry_creation() {
2189        let registry = ComponentRegistry::new();
2190        assert!(registry.list_tokenizers().is_empty());
2191    }
2192
2193    #[test]
2194    fn test_registry_with_defaults() {
2195        let registry = ComponentRegistry::with_defaults();
2196        assert!(registry.list_tokenizers().contains(&"stub".to_string()));
2197        assert!(registry
2198            .list_samplers()
2199            .contains(&"multinomial".to_string()));
2200        assert!(registry.list_schedulers().contains(&"fifo".to_string()));
2201    }
2202
2203    #[test]
2204    fn test_component_config() {
2205        let mut options = HashMap::new();
2206        options.insert("test_key".to_string(), serde_json::json!("test_value"));
2207
2208        let config = ComponentConfig {
2209            engine_config: EngineConfig::default(),
2210            device: Device::CPU,
2211            component_options: options,
2212            model_sources: None,
2213            prepared_model: None,
2214        };
2215
2216        assert_eq!(
2217            config.get_string_option("test_key"),
2218            Some("test_value".to_string())
2219        );
2220        assert!(config.get_string_option("missing").is_none());
2221    }
2222
2223    fn layer_split_component_config_for_device(device: usize) -> ComponentConfig {
2224        let mut config = EngineConfig::default();
2225        config.backend.device = Device::CUDA(device);
2226        config.backend.backend_options.insert(
2227            "selected_distributed_strategy".to_string(),
2228            serde_json::Value::String("layer_split".to_string()),
2229        );
2230        config.backend.backend_options.insert(
2231            "selected_gpu_devices".to_string(),
2232            serde_json::json!([0, 1]),
2233        );
2234        config.backend.backend_options.insert(
2235            "selected_layer_split_plan".to_string(),
2236            serde_json::Value::String(
2237                "stage0:cuda:0:layers=0-39;stage1:cuda:1:layers=40-79".to_string(),
2238            ),
2239        );
2240        config.backend.backend_options.insert(
2241            "selected_layer_split_stages".to_string(),
2242            serde_json::json!([
2243                {"stage": 0, "device": 0, "layer_start": 0, "layer_end": 39},
2244                {"stage": 1, "device": 1, "layer_start": 40, "layer_end": 79}
2245            ]),
2246        );
2247        ComponentConfig::from_engine_config(&config)
2248    }
2249
2250    #[test]
2251    fn resolves_llama_stage_config_for_selected_cuda_device() {
2252        let config = layer_split_component_config_for_device(1);
2253
2254        let stage = resolve_llama_layer_stage_config(&config, 80)
2255            .unwrap()
2256            .unwrap();
2257
2258        assert_eq!(stage.source_layers, 40..80);
2259        assert!(!stage.load_embedding);
2260        assert!(stage.load_lm_head);
2261    }
2262
2263    #[test]
2264    fn resolves_layer_split_pipeline_mode_default_and_explicit_batch() {
2265        let mut config = layer_split_component_config_for_device(0);
2266
2267        assert_eq!(
2268            resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2269            ferrum_models::models::LlamaPipelineMode::Overlapped
2270        );
2271
2272        config.component_options.insert(
2273            "layer_split_pipeline_mode".to_string(),
2274            serde_json::Value::String("batch".to_string()),
2275        );
2276        assert_eq!(
2277            resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2278            ferrum_models::models::LlamaPipelineMode::Batch
2279        );
2280    }
2281
2282    #[test]
2283    fn rejects_llama_stage_config_when_plan_layer_count_mismatches_model() {
2284        let config = layer_split_component_config_for_device(0);
2285
2286        let err = resolve_llama_layer_stage_config(&config, 81)
2287            .unwrap_err()
2288            .to_string();
2289
2290        assert!(err.contains("covers 80 layers but model has 81"));
2291    }
2292
2293    #[test]
2294    fn build_llm_rejects_layer_split_before_weight_open_without_device_scope() {
2295        use ferrum_interfaces::kv_dtype::KvFp16;
2296        use ferrum_models::models::LlamaFamilyConfig;
2297
2298        let qcfg = LlamaFamilyConfig {
2299            hidden_size: 1,
2300            intermediate_size: 1,
2301            num_heads: 1,
2302            num_kv_heads: 1,
2303            head_dim: 1,
2304            num_layers: 2,
2305            vocab_size: 1,
2306            max_seq_len: 8,
2307            rms_norm_eps: 1e-5,
2308            rope_theta: 10_000.0,
2309            rope_scaling: None,
2310            rope_interleaved: false,
2311            has_qk_norm: false,
2312            sliding_window: 0,
2313            ..Default::default()
2314        };
2315        let plan = crate::layer_split::parse_layer_split_plan(
2316            "stage0:cuda:0:layers=0-0;stage1:cuda:1:layers=1-1",
2317        )
2318        .unwrap();
2319
2320        let err = match build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
2321            ferrum_models::Architecture::Llama,
2322            qcfg,
2323            None,
2324            "/missing/model/path",
2325            Some(plan),
2326            Some(ferrum_models::models::LlamaPipelineMode::Overlapped),
2327        ) {
2328            Ok(_) => panic!("layer_split build unexpectedly succeeded"),
2329            Err(err) => err.to_string(),
2330        };
2331
2332        assert!(err.contains("device-scoped execution"));
2333        assert!(err.contains("refusing to silently use the default device"));
2334    }
2335
2336    #[test]
2337    fn test_registry_runtime_env_parses_model_dtype_and_tp() {
2338        let env = RegistryRuntimeEnv::from_env_vars([
2339            ("FERRUM_MODEL_PATH", "/models/qwen"),
2340            ("FERRUM_METAL_DTYPE", "fp16"),
2341            ("FERRUM_DTYPE", "fp32"),
2342            ("FERRUM_TP", "4"),
2343        ]);
2344
2345        assert_eq!(env.model_path.as_deref(), Some("/models/qwen"));
2346        assert_eq!(env.metal_dtype.as_deref(), Some("fp16"));
2347        assert_eq!(env.dtype.as_deref(), Some("fp32"));
2348        assert_eq!(env.tp, 4);
2349    }
2350
2351    #[test]
2352    fn test_registry_runtime_env_defaults_invalid_tp_and_cuda_dtype() {
2353        let env = RegistryRuntimeEnv::from_env_vars([
2354            ("FERRUM_DTYPE", "not-a-dtype"),
2355            ("FERRUM_TP", "not-a-number"),
2356        ]);
2357
2358        assert_eq!(env.model_path(), None);
2359        assert_eq!(env.tp, 0);
2360        assert_eq!(
2361            env.dtype_for_device(&Device::CUDA(0)),
2362            candle_core::DType::F16
2363        );
2364        assert_eq!(env.dtype_for_device(&Device::CPU), candle_core::DType::F32);
2365    }
2366
2367    #[cfg(any(target_os = "macos", target_os = "ios"))]
2368    #[test]
2369    fn test_registry_runtime_env_metal_dtype_precedence() {
2370        let env = RegistryRuntimeEnv::from_env_vars([
2371            ("FERRUM_METAL_DTYPE", "fp16"),
2372            ("FERRUM_DTYPE", "fp32"),
2373        ]);
2374
2375        assert_eq!(
2376            env.dtype_for_device(&Device::Metal),
2377            candle_core::DType::F16
2378        );
2379        assert_eq!(env.metal_dtype_hint(), "fp16");
2380    }
2381
2382    #[test]
2383    fn test_stub_tokenizer() {
2384        let tokenizer = StubTokenizer::new(100);
2385        assert_eq!(tokenizer.vocab_size(), 100);
2386
2387        let tokens = tokenizer.encode("hello world", false).unwrap();
2388        assert!(!tokens.is_empty());
2389
2390        let text = tokenizer.decode(&tokens, false).unwrap();
2391        assert!(text.contains("token_"));
2392    }
2393
2394    #[test]
2395    fn test_greedy_sampler() {
2396        use rand::SeedableRng;
2397
2398        let sampler = GreedySampler;
2399        let mut rng = rand::rngs::StdRng::seed_from_u64(42);
2400
2401        let logits = vec![0.1, 0.5, 0.3, 0.9, 0.2];
2402        let token = sampler.sample(&logits, &mut rng).unwrap();
2403
2404        assert_eq!(token.get(), 3); // Index of 0.9
2405    }
2406}