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}