everruns_contracts/
model_spec.rs1use serde::{Deserialize, Serialize};
4
5use crate::runtime_provider::{ProviderKey, RuntimeProviderRegistry};
6
7#[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(®istry)
98 .unwrap();
99 let west = ModelSpec::on("west", "shared-model")
100 .resolve_provider(®istry)
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(®istry)
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}