Skip to main content

ferrum_server/
model_registry.rs

1use ferrum_types::ModelId;
2use std::collections::HashMap;
3use std::fmt;
4
5#[derive(Clone, Copy, Debug, PartialEq, Eq)]
6pub enum ServedModelKind {
7    Llm,
8    Embedding,
9    Transcription,
10    Speech,
11}
12
13impl ServedModelKind {
14    pub fn modalities(self) -> &'static [&'static str] {
15        match self {
16            Self::Llm => &["text"],
17            Self::Embedding => &["text", "image"],
18            Self::Transcription => &["audio"],
19            Self::Speech => &["text", "audio"],
20        }
21    }
22}
23
24#[derive(Clone, Debug, PartialEq, Eq, Hash)]
25pub struct ServedModelName(String);
26
27impl ServedModelName {
28    pub fn as_str(&self) -> &str {
29        &self.0
30    }
31}
32
33impl fmt::Display for ServedModelName {
34    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
35        formatter.write_str(&self.0)
36    }
37}
38
39#[derive(Clone, Debug, PartialEq, Eq)]
40pub struct LoraAdapterModel {
41    pub name: String,
42    pub model_id: String,
43    pub path: String,
44}
45
46impl LoraAdapterModel {
47    pub fn new(
48        name: impl Into<String>,
49        model_id: impl Into<String>,
50        path: impl Into<String>,
51    ) -> Self {
52        Self {
53            name: name.into(),
54            model_id: model_id.into(),
55            path: path.into(),
56        }
57    }
58}
59
60#[derive(Clone, Debug)]
61pub struct ServedModelEntry {
62    public_name: ServedModelName,
63    engine_model_id: ModelId,
64    kind: ServedModelKind,
65    parent_public_name: Option<ServedModelName>,
66    adapter: Option<LoraAdapterModel>,
67}
68
69impl ServedModelEntry {
70    pub fn public_name(&self) -> &ServedModelName {
71        &self.public_name
72    }
73
74    pub fn engine_model_id(&self) -> &ModelId {
75        &self.engine_model_id
76    }
77
78    pub fn kind(&self) -> ServedModelKind {
79        self.kind
80    }
81
82    pub fn parent_public_name(&self) -> Option<&ServedModelName> {
83        self.parent_public_name.as_ref()
84    }
85
86    pub fn adapter(&self) -> Option<&LoraAdapterModel> {
87        self.adapter.as_ref()
88    }
89}
90
91#[derive(Clone, Debug, Default)]
92pub struct ServedModelRegistry {
93    entries: Vec<ServedModelEntry>,
94    entry_by_public_name: HashMap<String, usize>,
95}
96
97#[derive(Clone, Debug, PartialEq, Eq)]
98pub struct ServedModelRegistryError(String);
99
100impl fmt::Display for ServedModelRegistryError {
101    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
102        formatter.write_str(&self.0)
103    }
104}
105
106impl std::error::Error for ServedModelRegistryError {}
107
108impl ServedModelRegistry {
109    pub fn try_new(
110        engine_model_id: impl Into<ModelId>,
111        kind: ServedModelKind,
112        served_model_names: Vec<String>,
113        adapters: Vec<LoraAdapterModel>,
114    ) -> Result<Self, ServedModelRegistryError> {
115        let engine_model_id = engine_model_id.into();
116        validate_identifier("engine model id", &engine_model_id.0)?;
117        if served_model_names.is_empty() {
118            return Err(ServedModelRegistryError(
119                "at least one served model name is required".to_string(),
120            ));
121        }
122        if kind != ServedModelKind::Llm && !adapters.is_empty() {
123            return Err(ServedModelRegistryError(
124                "LoRA adapters may only be registered for LLM models".to_string(),
125            ));
126        }
127
128        let mut registry = Self {
129            entries: Vec::with_capacity(served_model_names.len() + adapters.len()),
130            entry_by_public_name: HashMap::with_capacity(served_model_names.len() + adapters.len()),
131        };
132        for public_name in served_model_names {
133            registry.push_entry(public_name, engine_model_id.clone(), kind, None, None)?;
134        }
135        let primary_public_name = registry.entries[0].public_name.clone();
136        let mut adapter_name_indexes = HashMap::with_capacity(adapters.len());
137        for adapter in adapters {
138            validate_identifier("LoRA adapter name", &adapter.name)?;
139            validate_identifier("LoRA public model id", &adapter.model_id)?;
140            if adapter.path.trim().is_empty() {
141                return Err(ServedModelRegistryError(format!(
142                    "LoRA adapter path must not be empty: {}",
143                    adapter.name
144                )));
145            }
146            if adapter_name_indexes
147                .insert(adapter.name.clone(), registry.entries.len())
148                .is_some()
149            {
150                return Err(ServedModelRegistryError(format!(
151                    "duplicate LoRA adapter name: {}",
152                    adapter.name
153                )));
154            }
155            registry.push_entry(
156                adapter.model_id.clone(),
157                engine_model_id.clone(),
158                kind,
159                Some(primary_public_name.clone()),
160                Some(adapter),
161            )?;
162        }
163        Ok(registry)
164    }
165
166    fn push_entry(
167        &mut self,
168        public_name: String,
169        engine_model_id: ModelId,
170        kind: ServedModelKind,
171        parent_public_name: Option<ServedModelName>,
172        adapter: Option<LoraAdapterModel>,
173    ) -> Result<(), ServedModelRegistryError> {
174        validate_identifier("served model name", &public_name)?;
175        if self.entry_by_public_name.contains_key(&public_name) {
176            return Err(ServedModelRegistryError(format!(
177                "duplicate or colliding served model name: {public_name}"
178            )));
179        }
180        let index = self.entries.len();
181        self.entry_by_public_name.insert(public_name.clone(), index);
182        self.entries.push(ServedModelEntry {
183            public_name: ServedModelName(public_name),
184            engine_model_id,
185            kind,
186            parent_public_name,
187            adapter,
188        });
189        Ok(())
190    }
191
192    pub fn entries(&self) -> &[ServedModelEntry] {
193        &self.entries
194    }
195
196    pub fn is_empty(&self) -> bool {
197        self.entries.is_empty()
198    }
199
200    pub fn primary_model_name(&self) -> Option<&ServedModelName> {
201        self.entries
202            .iter()
203            .find(|entry| entry.parent_public_name.is_none())
204            .map(|entry| &entry.public_name)
205    }
206
207    pub fn adapter_models(&self) -> impl Iterator<Item = &LoraAdapterModel> {
208        self.entries.iter().filter_map(ServedModelEntry::adapter)
209    }
210
211    pub fn adapter_count(&self) -> usize {
212        self.entries
213            .iter()
214            .filter(|entry| entry.adapter.is_some())
215            .count()
216    }
217
218    pub fn try_with_lora_adapters(
219        &self,
220        expected_base: &str,
221        adapters: Vec<LoraAdapterModel>,
222    ) -> Result<Self, ServedModelRegistryError> {
223        let base_entries = self
224            .entries
225            .iter()
226            .filter(|entry| entry.parent_public_name.is_none())
227            .collect::<Vec<_>>();
228        let primary = base_entries.first().ok_or_else(|| {
229            ServedModelRegistryError("cannot attach LoRA adapters to an empty registry".to_string())
230        })?;
231        if primary.kind != ServedModelKind::Llm {
232            return Err(ServedModelRegistryError(
233                "LoRA adapters may only be registered for LLM models".to_string(),
234            ));
235        }
236        require_same_base_model(&base_entries)?;
237        let expected_base_matches = primary.engine_model_id.0 == expected_base
238            || base_entries
239                .iter()
240                .any(|entry| entry.public_name.as_str() == expected_base);
241        if !expected_base_matches {
242            return Err(ServedModelRegistryError(format!(
243                "LoRA base {expected_base} does not match the registered engine or public model"
244            )));
245        }
246        Self::try_new(
247            primary.engine_model_id.clone(),
248            primary.kind,
249            base_entries
250                .iter()
251                .map(|entry| entry.public_name.to_string())
252                .collect(),
253            adapters,
254        )
255    }
256
257    pub fn resolve(
258        &self,
259        request_model: &str,
260        required_kind: ServedModelKind,
261    ) -> Option<&ServedModelEntry> {
262        let entry = self
263            .entry_by_public_name
264            .get(request_model)
265            .and_then(|index| self.entries.get(*index))?;
266        (entry.kind == required_kind).then_some(entry)
267    }
268}
269
270fn require_same_base_model(entries: &[&ServedModelEntry]) -> Result<(), ServedModelRegistryError> {
271    let Some(first) = entries.first() else {
272        return Ok(());
273    };
274    if entries
275        .iter()
276        .any(|entry| entry.engine_model_id != first.engine_model_id || entry.kind != first.kind)
277    {
278        return Err(ServedModelRegistryError(
279            "served model aliases do not resolve to one engine model and kind".to_string(),
280        ));
281    }
282    Ok(())
283}
284
285fn validate_identifier(label: &str, value: &str) -> Result<(), ServedModelRegistryError> {
286    if value.is_empty() || value.trim() != value {
287        return Err(ServedModelRegistryError(format!(
288            "{label} must be non-empty and have no surrounding whitespace"
289        )));
290    }
291    Ok(())
292}
293
294#[cfg(test)]
295mod tests {
296    use super::*;
297
298    #[test]
299    fn resolves_public_aliases_to_one_engine_model() {
300        let registry = ServedModelRegistry::try_new(
301            "internal-model",
302            ServedModelKind::Llm,
303            vec!["public-model".to_string(), "short-alias".to_string()],
304            vec![],
305        )
306        .unwrap();
307
308        for name in ["public-model", "short-alias"] {
309            let entry = registry.resolve(name, ServedModelKind::Llm).unwrap();
310            assert_eq!(entry.engine_model_id(), &ModelId::new("internal-model"));
311            assert!(entry.adapter().is_none());
312        }
313        assert!(registry
314            .resolve("internal-model", ServedModelKind::Llm)
315            .is_none());
316        assert!(registry
317            .resolve("public-model", ServedModelKind::Embedding)
318            .is_none());
319    }
320
321    #[test]
322    fn resolves_lora_model_to_internal_id_and_public_parent() {
323        let registry = ServedModelRegistry::try_new(
324            "internal-model",
325            ServedModelKind::Llm,
326            vec!["public-model".to_string()],
327            vec![LoraAdapterModel::new(
328                "sql",
329                "public-model:sql",
330                "/models/sql",
331            )],
332        )
333        .unwrap();
334
335        let entry = registry
336            .resolve("public-model:sql", ServedModelKind::Llm)
337            .unwrap();
338        assert_eq!(entry.engine_model_id(), &ModelId::new("internal-model"));
339        assert_eq!(entry.adapter().unwrap().name, "sql");
340        assert_eq!(
341            entry.parent_public_name().map(ServedModelName::as_str),
342            Some("public-model")
343        );
344    }
345
346    #[test]
347    fn adding_lora_preserves_all_public_base_aliases() {
348        let base = ServedModelRegistry::try_new(
349            "internal-model",
350            ServedModelKind::Llm,
351            vec!["public-model".to_string(), "short-alias".to_string()],
352            vec![],
353        )
354        .unwrap();
355        let registry = base
356            .try_with_lora_adapters(
357                "public-model",
358                vec![LoraAdapterModel::new(
359                    "sql",
360                    "public-model:sql",
361                    "/models/sql",
362                )],
363            )
364            .unwrap();
365
366        assert!(registry
367            .resolve("short-alias", ServedModelKind::Llm)
368            .is_some());
369        assert_eq!(
370            registry
371                .resolve("public-model:sql", ServedModelKind::Llm)
372                .unwrap()
373                .parent_public_name()
374                .map(ServedModelName::as_str),
375            Some("public-model")
376        );
377    }
378
379    #[test]
380    fn rejects_ambiguous_or_invalid_public_names() {
381        for result in [
382            ServedModelRegistry::try_new("internal", ServedModelKind::Llm, vec![], vec![]),
383            ServedModelRegistry::try_new(
384                "internal",
385                ServedModelKind::Llm,
386                vec!["same".to_string(), "same".to_string()],
387                vec![],
388            ),
389            ServedModelRegistry::try_new(
390                "internal",
391                ServedModelKind::Llm,
392                vec!["base".to_string()],
393                vec![LoraAdapterModel::new("sql", "base", "/models/sql")],
394            ),
395        ] {
396            assert!(result.is_err());
397        }
398    }
399}