Skip to main content

llm/
provider_connection.rs

1use std::collections::BTreeMap;
2use std::num::NonZeroU64;
3use std::time::Duration;
4
5pub(crate) const DEFAULT_STREAM_IDLE_TIMEOUT: Duration = Duration::from_mins(5);
6
7#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, serde::Deserialize, serde::Serialize, schemars::JsonSchema)]
8#[serde(rename_all = "kebab-case")]
9pub enum ProviderAuthMode {
10    #[default]
11    Default,
12    None,
13}
14
15#[derive(Clone, Debug, PartialEq, Eq)]
16pub struct ProviderConnectionConfig {
17    pub base_url: Option<String>,
18    pub auth_mode: ProviderAuthMode,
19    pub request_model: Option<String>,
20    pub inference_profile_arn: Option<String>,
21    pub idle_timeout: Duration,
22}
23
24#[doc = include_str!("docs/provider_connection_override.md")]
25#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Deserialize, serde::Serialize, schemars::JsonSchema)]
26#[serde(rename_all = "camelCase", deny_unknown_fields)]
27pub struct ProviderConnectionOverride {
28    /// Base URL override for the provider's API endpoint.
29    #[serde(default, rename = "url", skip_serializing_if = "Option::is_none")]
30    pub base_url: Option<String>,
31    /// Authentication mode. `default` uses the provider's normal credential
32    /// chain; `none` disables auth, for local or unauthenticated servers.
33    #[serde(default, rename = "auth", skip_serializing_if = "Option::is_none")]
34    pub auth_mode: Option<ProviderAuthMode>,
35    /// Provider-specific model or deployment target sent in requests without changing catalog identity.
36    #[serde(default, skip_serializing_if = "Option::is_none")]
37    pub request_model: Option<String>,
38    /// AWS Bedrock application inference profile ARN to route requests through.
39    #[serde(default, skip_serializing_if = "Option::is_none")]
40    pub inference_profile_arn: Option<String>,
41    #[serde(default, skip_serializing_if = "Option::is_none")]
42    pub idle_timeout_secs: Option<NonZeroU64>,
43}
44
45#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Deserialize, serde::Serialize, schemars::JsonSchema)]
46#[serde(transparent)]
47pub struct ProviderConnectionOverrides {
48    providers: BTreeMap<String, ProviderConnectionOverride>,
49}
50
51impl ProviderConnectionConfig {
52    pub fn from_override(value: ProviderConnectionOverride) -> Self {
53        Self {
54            base_url: value.base_url,
55            auth_mode: value.auth_mode.unwrap_or_default(),
56            request_model: value.request_model,
57            inference_profile_arn: value.inference_profile_arn,
58            idle_timeout: value
59                .idle_timeout_secs
60                .map_or(DEFAULT_STREAM_IDLE_TIMEOUT, |secs| Duration::from_secs(secs.get())),
61        }
62    }
63}
64
65impl Default for ProviderConnectionConfig {
66    fn default() -> Self {
67        Self {
68            base_url: None,
69            auth_mode: ProviderAuthMode::default(),
70            request_model: None,
71            inference_profile_arn: None,
72            idle_timeout: DEFAULT_STREAM_IDLE_TIMEOUT,
73        }
74    }
75}
76
77impl ProviderConnectionOverride {
78    pub fn url(url: impl Into<String>) -> Self {
79        Self { base_url: Some(url.into()), ..Self::default() }
80    }
81
82    pub fn auth(auth_mode: ProviderAuthMode) -> Self {
83        Self { auth_mode: Some(auth_mode), ..Self::default() }
84    }
85
86    pub fn request_model(model: impl Into<String>) -> Self {
87        Self { request_model: Some(model.into()), ..Self::default() }
88    }
89
90    pub fn inference_profile_arn(arn: impl Into<String>) -> Self {
91        Self { inference_profile_arn: Some(arn.into()), ..Self::default() }
92    }
93
94    pub fn idle_timeout_secs(secs: NonZeroU64) -> Self {
95        Self { idle_timeout_secs: Some(secs), ..Self::default() }
96    }
97
98    pub fn merge(&mut self, override_value: Self) {
99        if override_value.base_url.is_some() {
100            self.base_url = override_value.base_url;
101        }
102        if override_value.auth_mode.is_some() {
103            self.auth_mode = override_value.auth_mode;
104        }
105        if override_value.request_model.is_some() {
106            self.request_model = override_value.request_model;
107        }
108        if override_value.inference_profile_arn.is_some() {
109            self.inference_profile_arn = override_value.inference_profile_arn;
110        }
111        if override_value.idle_timeout_secs.is_some() {
112            self.idle_timeout_secs = override_value.idle_timeout_secs;
113        }
114    }
115}
116
117impl ProviderConnectionOverrides {
118    pub fn new(providers: BTreeMap<String, ProviderConnectionOverride>) -> Self {
119        Self { providers }
120    }
121
122    pub fn is_empty(&self) -> bool {
123        self.providers.is_empty()
124    }
125
126    pub fn merge(&mut self, overrides: ProviderConnectionOverrides) {
127        for (provider, override_value) in overrides.providers {
128            self.providers
129                .entry(provider)
130                .and_modify(|existing| existing.merge(override_value.clone()))
131                .or_insert(override_value);
132        }
133    }
134
135    pub fn config_for(&self, provider: &str) -> ProviderConnectionConfig {
136        self.providers.get(provider).cloned().map(ProviderConnectionConfig::from_override).unwrap_or_default()
137    }
138
139    pub fn into_inner(self) -> BTreeMap<String, ProviderConnectionOverride> {
140        self.providers
141    }
142}
143
144#[cfg(test)]
145mod tests {
146    use super::*;
147
148    #[test]
149    fn deserializes_and_merges_request_model() {
150        let mut first: ProviderConnectionOverrides =
151            serde_json::from_str(r#"{"azure-foundry":{"requestModel":"first"}}"#).unwrap();
152        first.merge(ProviderConnectionOverrides::new(BTreeMap::from([(
153            "azure-foundry".to_string(),
154            ProviderConnectionOverride::request_model("second"),
155        )])));
156
157        assert_eq!(first.config_for("azure-foundry").request_model.as_deref(), Some("second"));
158    }
159    #[test]
160    fn deserializes_bedrock_inference_profile_arn() {
161        let overrides: ProviderConnectionOverrides = serde_json::from_str(
162            r#"{"bedrock":{"inferenceProfileArn":"arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000"}}"#,
163        )
164        .unwrap();
165
166        let config = overrides.config_for("bedrock");
167
168        assert_eq!(
169            config.inference_profile_arn.as_deref(),
170            Some("arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000")
171        );
172    }
173
174    #[test]
175    fn idle_timeout_defaults_and_deserializes() {
176        let overrides: ProviderConnectionOverrides =
177            serde_json::from_str(r#"{"ollama":{"idleTimeoutSecs":900}}"#).unwrap();
178
179        assert_eq!(overrides.config_for("ollama").idle_timeout, Duration::from_mins(15));
180        assert_eq!(overrides.config_for("anthropic").idle_timeout, Duration::from_mins(5));
181        assert!(serde_json::from_str::<ProviderConnectionOverrides>(r#"{"ollama":{"idleTimeoutSecs":0}}"#).is_err());
182    }
183
184    #[test]
185    fn merge_replaces_inference_profile_arn() {
186        let mut first = ProviderConnectionOverride::inference_profile_arn("arn:first");
187
188        first.merge(ProviderConnectionOverride::inference_profile_arn("arn:second"));
189
190        assert_eq!(first.inference_profile_arn.as_deref(), Some("arn:second"));
191    }
192
193    #[test]
194    fn provider_overrides_merge_inference_profile_arn() {
195        let mut first = ProviderConnectionOverrides::new(BTreeMap::from([(
196            "bedrock".to_string(),
197            ProviderConnectionOverride::inference_profile_arn("arn:first"),
198        )]));
199        let second = ProviderConnectionOverrides::new(BTreeMap::from([(
200            "bedrock".to_string(),
201            ProviderConnectionOverride::inference_profile_arn("arn:second"),
202        )]));
203
204        first.merge(second);
205
206        assert_eq!(first.config_for("bedrock").inference_profile_arn.as_deref(), Some("arn:second"));
207    }
208}