Skip to main content

everruns_contracts/
model_spec.rs

1//! Credential-free model selection.
2
3use serde::{Deserialize, Serialize};
4
5use crate::runtime_provider::{ProviderKey, RuntimeProviderRegistry};
6
7/// A model name paired with the runtime provider that serves it.
8///
9/// This is ordinary serializable configuration: credentials, endpoints, and
10/// service authentication live exclusively on the runtime provider.
11#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
12#[non_exhaustive]
13pub struct ModelSpec {
14    pub provider: ProviderKey,
15    pub model: String,
16}
17
18impl ModelSpec {
19    pub fn on(provider: impl Into<ProviderKey>, model: impl Into<String>) -> Self {
20        Self {
21            provider: provider.into(),
22            model: model.into(),
23        }
24    }
25
26    pub fn resolve_provider(
27        &self,
28        registry: &RuntimeProviderRegistry,
29    ) -> Result<std::sync::Arc<crate::RuntimeProvider>, UnknownProvider> {
30        registry.get(&self.provider).ok_or_else(|| UnknownProvider {
31            requested: self.provider.clone(),
32            registered: registry.ids(),
33        })
34    }
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
38pub struct UnknownProvider {
39    pub requested: ProviderKey,
40    pub registered: Vec<String>,
41}
42
43impl std::fmt::Display for UnknownProvider {
44    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45        write!(
46            f,
47            "provider '{}' is not registered; registered providers: [{}]",
48            self.requested,
49            self.registered.join(", ")
50        )
51    }
52}
53
54impl std::error::Error for UnknownProvider {}
55
56#[cfg(test)]
57mod tests {
58    use super::*;
59
60    struct Noop;
61
62    #[async_trait::async_trait]
63    impl crate::ChatDriver for Noop {
64        async fn chat_completion_stream(
65            &self,
66            _endpoint: &crate::ProviderEndpoint,
67            _messages: Vec<crate::Message>,
68            _config: &crate::LlmCallConfig,
69        ) -> crate::Result<crate::LlmResponseStream> {
70            unreachable!()
71        }
72    }
73
74    #[test]
75    fn model_spec_is_credential_and_endpoint_free() {
76        let spec: ModelSpec = serde_json::from_str(r#"{"provider":" OpenAI-PROD ","model":"wire-model","api_key":"must-not-survive","base_url":"https://private.example"}"#).unwrap();
77        assert_eq!(spec, ModelSpec::on("openai-prod", "wire-model"));
78        assert_eq!(
79            serde_json::to_string(&spec).unwrap(),
80            r#"{"provider":"openai-prod","model":"wire-model"}"#
81        );
82        assert!(!format!("{spec:?}").contains("must-not-survive"));
83    }
84
85    #[test]
86    fn model_resolution_selects_provider_and_reports_sorted_alternatives() {
87        let driver: std::sync::Arc<dyn crate::ChatDriver> = std::sync::Arc::new(Noop);
88        let mut registry = RuntimeProviderRegistry::new();
89        registry
90            .register(crate::Provider::from_driver("east", driver.clone()))
91            .unwrap();
92        registry
93            .register(crate::Provider::from_driver("west", driver))
94            .unwrap();
95
96        let east = ModelSpec::on("east", "shared-model")
97            .resolve_provider(&registry)
98            .unwrap();
99        let west = ModelSpec::on("west", "shared-model")
100            .resolve_provider(&registry)
101            .unwrap();
102        assert_eq!(east.id().as_str(), "east");
103        assert_eq!(west.id().as_str(), "west");
104        assert!(std::sync::Arc::ptr_eq(
105            east.driver().unwrap(),
106            west.driver().unwrap()
107        ));
108        let error = ModelSpec::on(" Missing ", "shared-model")
109            .resolve_provider(&registry)
110            .unwrap_err();
111        assert_eq!(
112            error,
113            UnknownProvider {
114                requested: ProviderKey::new("missing"),
115                registered: vec!["east".into(), "west".into()]
116            }
117        );
118        assert_eq!(
119            error.to_string(),
120            "provider 'missing' is not registered; registered providers: [east, west]"
121        );
122        let empty = ModelSpec::on("missing", "model")
123            .resolve_provider(&RuntimeProviderRegistry::new())
124            .unwrap_err();
125        assert_eq!(
126            empty.to_string(),
127            "provider 'missing' is not registered; registered providers: []"
128        );
129    }
130}