llm/
provider_connection.rs1use 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 #[serde(default, rename = "url", skip_serializing_if = "Option::is_none")]
30 pub base_url: Option<String>,
31 #[serde(default, rename = "auth", skip_serializing_if = "Option::is_none")]
34 pub auth_mode: Option<ProviderAuthMode>,
35 #[serde(default, skip_serializing_if = "Option::is_none")]
37 pub request_model: Option<String>,
38 #[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}