Skip to main content

tokio_prompt_orchestrator/
model_registry.rs

1//! # Model Registry
2//!
3//! Tracks LLM model metadata, capabilities, pricing, and rate limits.
4//! Provides routing hints and deprecation warnings.
5
6use dashmap::DashMap;
7use std::fmt;
8
9// ── ModelCapability ───────────────────────────────────────────────────────────
10
11/// A capability a model may support.
12#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13pub enum ModelCapability {
14    /// Free-form text generation.
15    TextGeneration,
16    /// Code generation and completion.
17    CodeGeneration,
18    /// Structured function / tool calling.
19    FunctionCalling,
20    /// Accepts image inputs.
21    ImageUnderstanding,
22    /// Produces dense vector embeddings.
23    Embedding,
24    /// Condensed document summarization.
25    Summarization,
26    /// Language translation.
27    Translation,
28    /// Multi-class text classification.
29    Classification,
30}
31
32impl fmt::Display for ModelCapability {
33    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
34        let s = match self {
35            Self::TextGeneration => "TextGeneration",
36            Self::CodeGeneration => "CodeGeneration",
37            Self::FunctionCalling => "FunctionCalling",
38            Self::ImageUnderstanding => "ImageUnderstanding",
39            Self::Embedding => "Embedding",
40            Self::Summarization => "Summarization",
41            Self::Translation => "Translation",
42            Self::Classification => "Classification",
43        };
44        write!(f, "{}", s)
45    }
46}
47
48// ── ModelTier ─────────────────────────────────────────────────────────────────
49
50/// Tier classification for a model, reflecting quality and price.
51#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
52pub enum ModelTier {
53    /// Cheapest, lowest capability.
54    Economy,
55    /// Balanced cost and capability.
56    Standard,
57    /// High capability, higher cost.
58    Advanced,
59    /// Flagship quality.
60    Premium,
61}
62
63impl ModelTier {
64    /// Relative pricing weight for budget estimation (Economy=1.0 … Premium=8.0).
65    pub fn pricing_weight(&self) -> f64 {
66        match self {
67            Self::Economy => 1.0,
68            Self::Standard => 2.5,
69            Self::Advanced => 5.0,
70            Self::Premium => 8.0,
71        }
72    }
73}
74
75impl fmt::Display for ModelTier {
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        let s = match self {
78            Self::Economy => "Economy",
79            Self::Standard => "Standard",
80            Self::Advanced => "Advanced",
81            Self::Premium => "Premium",
82        };
83        write!(f, "{}", s)
84    }
85}
86
87// ── ModelMetadata ─────────────────────────────────────────────────────────────
88
89/// Rich metadata for a single LLM model.
90#[derive(Debug, Clone)]
91pub struct ModelMetadata {
92    /// Unique model identifier (e.g. "gpt-4o").
93    pub id: String,
94    /// Human-readable display name.
95    pub name: String,
96    /// Provider name (e.g. "OpenAI").
97    pub provider: String,
98    /// Quality/price tier.
99    pub tier: ModelTier,
100    /// Maximum context window in tokens.
101    pub context_window: usize,
102    /// Maximum tokens the model can output in one call.
103    pub max_output_tokens: usize,
104    /// Supported capabilities.
105    pub capabilities: Vec<ModelCapability>,
106    /// Cost per 1 000 input tokens in USD.
107    pub cost_per_1k_input_tokens: f64,
108    /// Cost per 1 000 output tokens in USD.
109    pub cost_per_1k_output_tokens: f64,
110    /// API rate limit: requests per minute.
111    pub requests_per_minute: u32,
112    /// API rate limit: tokens per minute.
113    pub tokens_per_minute: u32,
114    /// Whether the model supports streaming responses.
115    pub supports_streaming: bool,
116    /// Whether the model accepts a system prompt.
117    pub supports_system_prompt: bool,
118    /// Whether this model is deprecated.
119    pub deprecated: bool,
120    /// ID of the recommended successor if deprecated.
121    pub successor: Option<String>,
122}
123
124// ── RoutingHint ───────────────────────────────────────────────────────────────
125
126/// Routing suggestion produced by [`ModelRegistry::suggest_routing`].
127#[derive(Debug, Clone)]
128pub struct RoutingHint {
129    /// Recommended model ID.
130    pub preferred_model: String,
131    /// Ordered list of fallback model IDs.
132    pub fallback_models: Vec<String>,
133    /// Human-readable rationale.
134    pub reason: String,
135    /// Confidence in [0.0, 1.0].
136    pub confidence: f64,
137}
138
139// ── ModelRegistry ─────────────────────────────────────────────────────────────
140
141/// Concurrent registry of [`ModelMetadata`] keyed by model ID.
142pub struct ModelRegistry {
143    models: DashMap<String, ModelMetadata>,
144}
145
146impl Default for ModelRegistry {
147    fn default() -> Self {
148        Self {
149            models: DashMap::new(),
150        }
151    }
152}
153
154impl ModelRegistry {
155    /// Create an empty registry.
156    pub fn new() -> Self {
157        Self::default()
158    }
159
160    /// Register or replace a model's metadata.
161    pub fn register(&self, model: ModelMetadata) {
162        self.models.insert(model.id.clone(), model);
163    }
164
165    /// Look up a model by ID.
166    pub fn get(&self, model_id: &str) -> Option<ModelMetadata> {
167        self.models.get(model_id).map(|r| r.clone())
168    }
169
170    /// Return all registered models.
171    pub fn all_models(&self) -> Vec<ModelMetadata> {
172        self.models.iter().map(|r| r.clone()).collect()
173    }
174
175    /// Return non-deprecated models that have `cap`, sorted by cheapest input cost first.
176    pub fn models_with_capability(&self, cap: &ModelCapability) -> Vec<ModelMetadata> {
177        let mut models: Vec<ModelMetadata> = self
178            .models
179            .iter()
180            .filter(|r| !r.deprecated && r.capabilities.contains(cap))
181            .map(|r| r.clone())
182            .collect();
183        models.sort_by(|a, b| {
184            a.cost_per_1k_input_tokens
185                .partial_cmp(&b.cost_per_1k_input_tokens)
186                .unwrap_or(std::cmp::Ordering::Equal)
187        });
188        models
189    }
190
191    /// Return the cheapest model that has `cap` and fits within `min_context` tokens.
192    pub fn cheapest_for_capability(
193        &self,
194        cap: &ModelCapability,
195        min_context: usize,
196    ) -> Option<ModelMetadata> {
197        self.models_with_capability(cap)
198            .into_iter()
199            .find(|m| m.context_window >= min_context)
200    }
201
202    /// Return the highest-tier model that has `cap` and fits within `min_context` tokens.
203    pub fn best_for_capability(
204        &self,
205        cap: &ModelCapability,
206        min_context: usize,
207    ) -> Option<ModelMetadata> {
208        let mut candidates: Vec<ModelMetadata> = self
209            .models_with_capability(cap)
210            .into_iter()
211            .filter(|m| m.context_window >= min_context)
212            .collect();
213        candidates.sort_by(|a, b| b.tier.cmp(&a.tier));
214        candidates.into_iter().next()
215    }
216
217    /// Suggest a routing strategy given required capabilities, context size, and optional budget.
218    pub fn suggest_routing(
219        &self,
220        caps: &[ModelCapability],
221        context_size: usize,
222        budget_usd: Option<f64>,
223    ) -> RoutingHint {
224        // Collect candidates satisfying all caps.
225        let mut candidates: Vec<ModelMetadata> = self
226            .models
227            .iter()
228            .filter(|r| {
229                !r.deprecated
230                    && r.context_window >= context_size
231                    && caps.iter().all(|c| r.capabilities.contains(c))
232            })
233            .map(|r| r.clone())
234            .collect();
235
236        if candidates.is_empty() {
237            return RoutingHint {
238                preferred_model: String::new(),
239                fallback_models: vec![],
240                reason: "No models satisfy all requested capabilities".to_string(),
241                confidence: 0.0,
242            };
243        }
244
245        // Apply budget filter if provided.
246        if let Some(budget) = budget_usd {
247            let budget_filtered: Vec<ModelMetadata> = candidates
248                .iter()
249                .filter(|m| m.cost_per_1k_input_tokens * (context_size as f64 / 1000.0) <= budget)
250                .cloned()
251                .collect();
252            if !budget_filtered.is_empty() {
253                candidates = budget_filtered;
254            }
255        }
256
257        // Sort: cheapest first for economy routing; best tier for quality routing.
258        candidates.sort_by(|a, b| {
259            a.cost_per_1k_input_tokens
260                .partial_cmp(&b.cost_per_1k_input_tokens)
261                .unwrap_or(std::cmp::Ordering::Equal)
262        });
263
264        let preferred = candidates.remove(0);
265        let fallback_models: Vec<String> = candidates.iter().take(3).map(|m| m.id.clone()).collect();
266        let reason = format!(
267            "Selected {} (tier={}, cost=${:.4}/1k input) for capabilities: {}",
268            preferred.id,
269            preferred.tier,
270            preferred.cost_per_1k_input_tokens,
271            caps.iter()
272                .map(|c| c.to_string())
273                .collect::<Vec<_>>()
274                .join(", ")
275        );
276        let confidence = if fallback_models.is_empty() { 0.7 } else { 0.9 };
277
278        RoutingHint {
279            preferred_model: preferred.id,
280            fallback_models,
281            reason,
282            confidence,
283        }
284    }
285
286    /// Build a registry pre-populated with well-known models.
287    pub fn built_in_registry() -> Self {
288        let registry = Self::new();
289
290        // GPT-4o
291        registry.register(ModelMetadata {
292            id: "gpt-4o".to_string(),
293            name: "GPT-4o".to_string(),
294            provider: "OpenAI".to_string(),
295            tier: ModelTier::Premium,
296            context_window: 128_000,
297            max_output_tokens: 4_096,
298            capabilities: vec![
299                ModelCapability::TextGeneration,
300                ModelCapability::CodeGeneration,
301                ModelCapability::FunctionCalling,
302                ModelCapability::ImageUnderstanding,
303                ModelCapability::Summarization,
304                ModelCapability::Translation,
305                ModelCapability::Classification,
306            ],
307            cost_per_1k_input_tokens: 0.0025,
308            cost_per_1k_output_tokens: 0.01,
309            requests_per_minute: 10_000,
310            tokens_per_minute: 2_000_000,
311            supports_streaming: true,
312            supports_system_prompt: true,
313            deprecated: false,
314            successor: None,
315        });
316
317        // GPT-4o-mini
318        registry.register(ModelMetadata {
319            id: "gpt-4o-mini".to_string(),
320            name: "GPT-4o Mini".to_string(),
321            provider: "OpenAI".to_string(),
322            tier: ModelTier::Economy,
323            context_window: 128_000,
324            max_output_tokens: 16_384,
325            capabilities: vec![
326                ModelCapability::TextGeneration,
327                ModelCapability::CodeGeneration,
328                ModelCapability::FunctionCalling,
329                ModelCapability::Summarization,
330                ModelCapability::Translation,
331                ModelCapability::Classification,
332            ],
333            cost_per_1k_input_tokens: 0.000150,
334            cost_per_1k_output_tokens: 0.000600,
335            requests_per_minute: 30_000,
336            tokens_per_minute: 10_000_000,
337            supports_streaming: true,
338            supports_system_prompt: true,
339            deprecated: false,
340            successor: None,
341        });
342
343        // Claude 3.5 Sonnet
344        registry.register(ModelMetadata {
345            id: "claude-3-5-sonnet".to_string(),
346            name: "Claude 3.5 Sonnet".to_string(),
347            provider: "Anthropic".to_string(),
348            tier: ModelTier::Advanced,
349            context_window: 200_000,
350            max_output_tokens: 8_192,
351            capabilities: vec![
352                ModelCapability::TextGeneration,
353                ModelCapability::CodeGeneration,
354                ModelCapability::FunctionCalling,
355                ModelCapability::ImageUnderstanding,
356                ModelCapability::Summarization,
357                ModelCapability::Translation,
358                ModelCapability::Classification,
359            ],
360            cost_per_1k_input_tokens: 0.003,
361            cost_per_1k_output_tokens: 0.015,
362            requests_per_minute: 4_000,
363            tokens_per_minute: 800_000,
364            supports_streaming: true,
365            supports_system_prompt: true,
366            deprecated: false,
367            successor: None,
368        });
369
370        // Claude 3 Haiku
371        registry.register(ModelMetadata {
372            id: "claude-3-haiku".to_string(),
373            name: "Claude 3 Haiku".to_string(),
374            provider: "Anthropic".to_string(),
375            tier: ModelTier::Economy,
376            context_window: 200_000,
377            max_output_tokens: 4_096,
378            capabilities: vec![
379                ModelCapability::TextGeneration,
380                ModelCapability::CodeGeneration,
381                ModelCapability::Summarization,
382                ModelCapability::Translation,
383                ModelCapability::Classification,
384            ],
385            cost_per_1k_input_tokens: 0.00025,
386            cost_per_1k_output_tokens: 0.00125,
387            requests_per_minute: 4_000,
388            tokens_per_minute: 800_000,
389            supports_streaming: true,
390            supports_system_prompt: true,
391            deprecated: false,
392            successor: None,
393        });
394
395        // Gemini 1.5 Pro
396        registry.register(ModelMetadata {
397            id: "gemini-1.5-pro".to_string(),
398            name: "Gemini 1.5 Pro".to_string(),
399            provider: "Google".to_string(),
400            tier: ModelTier::Premium,
401            context_window: 1_048_576,
402            max_output_tokens: 8_192,
403            capabilities: vec![
404                ModelCapability::TextGeneration,
405                ModelCapability::CodeGeneration,
406                ModelCapability::FunctionCalling,
407                ModelCapability::ImageUnderstanding,
408                ModelCapability::Summarization,
409                ModelCapability::Translation,
410                ModelCapability::Classification,
411            ],
412            cost_per_1k_input_tokens: 0.00125,
413            cost_per_1k_output_tokens: 0.005,
414            requests_per_minute: 360,
415            tokens_per_minute: 4_000_000,
416            supports_streaming: true,
417            supports_system_prompt: true,
418            deprecated: false,
419            successor: None,
420        });
421
422        // Gemini 1.5 Flash
423        registry.register(ModelMetadata {
424            id: "gemini-1.5-flash".to_string(),
425            name: "Gemini 1.5 Flash".to_string(),
426            provider: "Google".to_string(),
427            tier: ModelTier::Standard,
428            context_window: 1_048_576,
429            max_output_tokens: 8_192,
430            capabilities: vec![
431                ModelCapability::TextGeneration,
432                ModelCapability::CodeGeneration,
433                ModelCapability::ImageUnderstanding,
434                ModelCapability::Summarization,
435                ModelCapability::Translation,
436                ModelCapability::Classification,
437            ],
438            cost_per_1k_input_tokens: 0.000075,
439            cost_per_1k_output_tokens: 0.0003,
440            requests_per_minute: 1_000,
441            tokens_per_minute: 4_000_000,
442            supports_streaming: true,
443            supports_system_prompt: true,
444            deprecated: false,
445            successor: None,
446        });
447
448        // Llama 3 70B
449        registry.register(ModelMetadata {
450            id: "llama-3-70b".to_string(),
451            name: "Llama 3 70B".to_string(),
452            provider: "Meta".to_string(),
453            tier: ModelTier::Advanced,
454            context_window: 8_192,
455            max_output_tokens: 4_096,
456            capabilities: vec![
457                ModelCapability::TextGeneration,
458                ModelCapability::CodeGeneration,
459                ModelCapability::Summarization,
460                ModelCapability::Translation,
461                ModelCapability::Classification,
462            ],
463            cost_per_1k_input_tokens: 0.00059,
464            cost_per_1k_output_tokens: 0.00079,
465            requests_per_minute: 6_000,
466            tokens_per_minute: 800_000,
467            supports_streaming: true,
468            supports_system_prompt: true,
469            deprecated: false,
470            successor: None,
471        });
472
473        // Llama 3 8B
474        registry.register(ModelMetadata {
475            id: "llama-3-8b".to_string(),
476            name: "Llama 3 8B".to_string(),
477            provider: "Meta".to_string(),
478            tier: ModelTier::Economy,
479            context_window: 8_192,
480            max_output_tokens: 4_096,
481            capabilities: vec![
482                ModelCapability::TextGeneration,
483                ModelCapability::CodeGeneration,
484                ModelCapability::Summarization,
485                ModelCapability::Classification,
486            ],
487            cost_per_1k_input_tokens: 0.00010,
488            cost_per_1k_output_tokens: 0.00010,
489            requests_per_minute: 6_000,
490            tokens_per_minute: 800_000,
491            supports_streaming: true,
492            supports_system_prompt: true,
493            deprecated: false,
494            successor: None,
495        });
496
497        registry
498    }
499
500    /// Check whether a model is within its rate limits.
501    ///
502    /// Returns `true` if the model can accept more requests/tokens, `false` if a limit is exceeded.
503    pub fn rate_limit_check(
504        &self,
505        model_id: &str,
506        requests_last_minute: u32,
507        tokens_last_minute: u32,
508    ) -> bool {
509        match self.get(model_id) {
510            Some(m) => {
511                requests_last_minute < m.requests_per_minute
512                    && tokens_last_minute < m.tokens_per_minute
513            }
514            None => false,
515        }
516    }
517
518    /// Return deprecation warning strings for all deprecated models.
519    pub fn deprecation_warnings(&self) -> Vec<String> {
520        self.models
521            .iter()
522            .filter(|r| r.deprecated)
523            .map(|r| {
524                let successor_hint = r
525                    .successor
526                    .as_deref()
527                    .map(|s| format!(" Please migrate to '{}'.", s))
528                    .unwrap_or_default();
529                format!(
530                    "Model '{}' ({}) is deprecated.{}",
531                    r.id, r.name, successor_hint
532                )
533            })
534            .collect()
535    }
536}
537
538#[cfg(test)]
539mod tests {
540    use super::*;
541
542    #[test]
543    fn built_in_has_gpt4o() {
544        let r = ModelRegistry::built_in_registry();
545        assert!(r.get("gpt-4o").is_some());
546    }
547
548    #[test]
549    fn models_with_capability_sorted_by_cost() {
550        let r = ModelRegistry::built_in_registry();
551        let models = r.models_with_capability(&ModelCapability::TextGeneration);
552        assert!(!models.is_empty());
553        for i in 1..models.len() {
554            assert!(
555                models[i].cost_per_1k_input_tokens >= models[i - 1].cost_per_1k_input_tokens
556            );
557        }
558    }
559
560    #[test]
561    fn cheapest_for_capability_returns_some() {
562        let r = ModelRegistry::built_in_registry();
563        let m = r.cheapest_for_capability(&ModelCapability::FunctionCalling, 0);
564        assert!(m.is_some());
565    }
566
567    #[test]
568    fn best_for_capability_returns_highest_tier() {
569        let r = ModelRegistry::built_in_registry();
570        let m = r.best_for_capability(&ModelCapability::ImageUnderstanding, 0);
571        assert!(m.is_some());
572        let m = m.unwrap();
573        assert!(m.tier == ModelTier::Premium || m.tier == ModelTier::Advanced);
574    }
575
576    #[test]
577    fn rate_limit_check_within_limit() {
578        let r = ModelRegistry::built_in_registry();
579        assert!(r.rate_limit_check("gpt-4o", 100, 10_000));
580    }
581
582    #[test]
583    fn rate_limit_check_exceeds_limit() {
584        let r = ModelRegistry::built_in_registry();
585        assert!(!r.rate_limit_check("gpt-4o", 999_999, 999_999_999));
586    }
587
588    #[test]
589    fn no_deprecation_warnings_in_built_in() {
590        let r = ModelRegistry::built_in_registry();
591        assert!(r.deprecation_warnings().is_empty());
592    }
593
594    #[test]
595    fn deprecation_warning_for_deprecated_model() {
596        let r = ModelRegistry::new();
597        r.register(ModelMetadata {
598            id: "old-model".to_string(),
599            name: "Old Model".to_string(),
600            provider: "Acme".to_string(),
601            tier: ModelTier::Standard,
602            context_window: 4_096,
603            max_output_tokens: 1_024,
604            capabilities: vec![ModelCapability::TextGeneration],
605            cost_per_1k_input_tokens: 0.01,
606            cost_per_1k_output_tokens: 0.02,
607            requests_per_minute: 100,
608            tokens_per_minute: 50_000,
609            supports_streaming: false,
610            supports_system_prompt: false,
611            deprecated: true,
612            successor: Some("new-model".to_string()),
613        });
614        let warnings = r.deprecation_warnings();
615        assert_eq!(warnings.len(), 1);
616        assert!(warnings[0].contains("old-model"));
617        assert!(warnings[0].contains("new-model"));
618    }
619
620    #[test]
621    fn suggest_routing_no_cap_match_returns_empty() {
622        let r = ModelRegistry::new(); // empty
623        let hint = r.suggest_routing(&[ModelCapability::Embedding], 0, None);
624        assert!(hint.preferred_model.is_empty());
625        assert!((hint.confidence - 0.0).abs() < f64::EPSILON);
626    }
627}