1use schemars::JsonSchema;
2use serde::{Deserialize, Serialize};
3use std::collections::{BTreeMap, BTreeSet};
4
5use super::Settings;
6
7#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
8#[serde(rename_all = "kebab-case")]
9pub enum CustomReasoningProtocol {
10 #[default]
11 GptLike,
12 AnthropicLike,
13}
14
15pub(crate) const MAX_CUSTOM_PROVIDER_REQUEST_HEADERS: usize = 32;
16
17#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
18#[serde(tag = "source", rename_all = "snake_case")]
19pub enum CustomProviderHeaderValue {
20 ConversationId,
21}
22
23#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
24pub struct CustomProviderConfig {
25 pub label: String,
26 pub base_url: String,
27 #[serde(
28 default,
29 deserialize_with = "deserialize_optional_env_var",
30 skip_serializing_if = "Option::is_none"
31 )]
32 pub api_key_env_var: Option<String>,
33 #[serde(
34 default,
35 deserialize_with = "deserialize_optional_models_dev_provider",
36 skip_serializing_if = "Option::is_none"
37 )]
38 pub models_dev_provider: Option<String>,
39 #[serde(
40 default,
41 deserialize_with = "deserialize_optional_fast_mode",
42 skip_serializing_if = "Option::is_none"
43 )]
44 pub fast_mode: Option<CustomProviderFastMode>,
45 #[serde(default = "default_use_responses_endpoint")]
46 pub use_responses_endpoint: bool,
47 #[serde(default, skip_serializing_if = "is_false")]
48 pub supports_text_verbosity: bool,
49 #[serde(default, skip_serializing_if = "is_gpt_like")]
50 pub reasoning_protocol: CustomReasoningProtocol,
51 #[serde(
52 default,
53 deserialize_with = "deserialize_extra_models",
54 skip_serializing_if = "Vec::is_empty"
55 )]
56 pub extra_models: Vec<String>,
57 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
58 pub request_headers: BTreeMap<String, CustomProviderHeaderValue>,
59}
60
61fn default_use_responses_endpoint() -> bool {
62 true
63}
64
65pub(super) fn validate_custom_provider_settings(settings: &Settings) -> anyhow::Result<()> {
66 for (id, custom) in &settings.custom_providers {
67 validate_custom_provider_id(id).map_err(|error| {
68 anyhow::anyhow!("custom provider '{id}' has invalid provider id: {error}")
69 })?;
70 validate_custom_provider_label(&custom.label).map_err(|error| {
71 anyhow::anyhow!("custom provider '{id}' has invalid label: {error}")
72 })?;
73 normalize_custom_provider_base_url(&custom.base_url).map_err(|error| {
74 anyhow::anyhow!("custom provider '{id}' has invalid base_url: {error}")
75 })?;
76 if let Some(env_var) = &custom.api_key_env_var {
77 validate_env_var_name(env_var).map_err(|error| {
78 anyhow::anyhow!("custom provider '{id}' has invalid api_key_env_var: {error}")
79 })?;
80 }
81 if let Some(fast_mode) = &custom.fast_mode {
82 validate_custom_provider_fast_mode(fast_mode).map_err(|error| {
83 anyhow::anyhow!("custom provider '{id}' has invalid fast_mode: {error}")
84 })?;
85 }
86 if let Some(models_dev_provider) = &custom.models_dev_provider {
87 validate_models_dev_provider_namespace(models_dev_provider).map_err(|error| {
88 anyhow::anyhow!("custom provider '{id}' has invalid models_dev_provider: {error}")
89 })?;
90 }
91 normalized_extra_models(&custom.extra_models).map_err(|error| {
92 anyhow::anyhow!("custom provider '{id}' has invalid extra_models: {error}")
93 })?;
94 validate_custom_provider_request_headers(id, &custom.request_headers)?;
95 }
96 Ok(())
97}
98
99fn validate_custom_provider_request_headers(
100 provider_id: &str,
101 headers: &BTreeMap<String, CustomProviderHeaderValue>,
102) -> anyhow::Result<()> {
103 if headers.len() > MAX_CUSTOM_PROVIDER_REQUEST_HEADERS {
104 anyhow::bail!(
105 "custom provider '{provider_id}' request_headers must contain at most {MAX_CUSTOM_PROVIDER_REQUEST_HEADERS} entries"
106 );
107 }
108 let mut normalized = BTreeSet::new();
109 for name in headers.keys() {
110 reqwest::header::HeaderName::from_bytes(name.as_bytes()).map_err(|_| {
111 anyhow::anyhow!(
112 "custom provider '{provider_id}' has invalid request header name '{name}'"
113 )
114 })?;
115 let lower = name.to_ascii_lowercase();
116 if !normalized.insert(lower.clone()) {
117 anyhow::bail!(
118 "custom provider '{provider_id}' has duplicate request header name '{name}' (names are case-insensitive)"
119 );
120 }
121 if matches!(
122 lower.as_str(),
123 "accept"
124 | "authorization"
125 | "content-length"
126 | "content-type"
127 | "host"
128 | "proxy-authorization"
129 | "transfer-encoding"
130 | "user-agent"
131 ) {
132 anyhow::bail!(
133 "custom provider '{provider_id}' request_headers must not override transport-owned header '{name}'"
134 );
135 }
136 }
137 Ok(())
138}
139
140fn validate_custom_provider_label(label: &str) -> anyhow::Result<()> {
141 let label = label.trim();
142 if label.is_empty() || label.len() > 100 {
143 anyhow::bail!("custom provider label must be non-empty and at most 100 characters");
144 }
145 if looks_like_secret_label(label) {
146 anyhow::bail!("custom provider label must not look like a secret value");
147 }
148 Ok(())
149}
150
151fn looks_like_secret_label(value: &str) -> bool {
152 let value = value.trim();
153 value.starts_with("sk-")
154 || value.starts_with("Bearer ")
155 || value.contains('=')
156 || (value.len() >= 48
157 && value
158 .chars()
159 .filter(|ch| ch.is_ascii_alphanumeric())
160 .count()
161 >= 40)
162}
163
164pub(crate) fn validate_custom_provider_id(id: &str) -> anyhow::Result<String> {
165 let id = id.trim();
166 if matches!(
167 id,
168 crate::providers::OPENAI_CODEX_PROVIDER
169 | crate::providers::ANTHROPIC_PROVIDER
170 | crate::providers::CLAUDE_SUBSCRIPTION_PROVIDER
171 | "claude-code"
172 ) {
173 anyhow::bail!("custom provider id '{id}' is reserved");
174 }
175 if id.len() > 63
176 || id.is_empty()
177 || !id.as_bytes()[0].is_ascii_lowercase()
178 || id.ends_with('-')
179 || !id
180 .chars()
181 .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
182 {
183 anyhow::bail!(
184 "custom provider id must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
185 );
186 }
187 Ok(id.to_string())
188}
189
190pub(crate) fn looks_like_secret_value(value: &str) -> bool {
191 let value = value.trim();
192 value.starts_with("sk-")
193 || value.starts_with("Bearer ")
194 || value.contains('=')
195 || value.chars().any(char::is_whitespace)
196 || (value.len() >= 48
197 && value
198 .chars()
199 .filter(|ch| ch.is_ascii_alphanumeric())
200 .count()
201 >= 40)
202}
203
204pub(crate) fn validate_env_var_name(name: &str) -> anyhow::Result<String> {
205 let name = name.trim();
206 if looks_like_secret_value(name) {
207 anyhow::bail!(
208 "API key environment variable name looks like a secret value; enter a variable name such as CUSTOM_PROVIDER_API_KEY"
209 );
210 }
211 if name.is_empty()
212 || !(name.as_bytes()[0].is_ascii_uppercase() || name.as_bytes()[0] == b'_')
213 || !name
214 .chars()
215 .all(|ch| ch.is_ascii_uppercase() || ch.is_ascii_digit() || ch == '_')
216 {
217 anyhow::bail!("API key environment variable name must match ^[A-Z_][A-Z0-9_]*$");
218 }
219 Ok(name.to_string())
220}
221
222pub(crate) fn validate_optional_env_var_name(name: &str) -> anyhow::Result<Option<String>> {
223 if name.trim().is_empty() {
224 return Ok(None);
225 }
226 validate_env_var_name(name).map(Some)
227}
228
229pub(crate) fn normalized_extra_models(extra_models: &[String]) -> anyhow::Result<Vec<String>> {
230 if extra_models.len() > 64 {
231 anyhow::bail!("extra_models must contain at most 64 model ids");
232 }
233 let mut seen = BTreeSet::new();
234 let mut normalized = Vec::new();
235 for model in extra_models {
236 let model = model.trim();
237 if model.is_empty() {
238 anyhow::bail!("extra_models entries must be non-empty");
239 }
240 if model.len() > 200 {
241 anyhow::bail!("extra_models entries must be at most 200 bytes");
242 }
243 if model
244 .chars()
245 .any(|ch| ch.is_ascii_control() || ch.is_ascii_whitespace())
246 {
247 anyhow::bail!(
248 "extra_models entries must not contain ASCII control characters or whitespace"
249 );
250 }
251 if looks_like_secret_value(model) {
252 anyhow::bail!("extra_models entries must not look like secret values");
253 }
254 if seen.insert(model.to_string()) {
255 normalized.push(model.to_string());
256 }
257 }
258 Ok(normalized)
259}
260
261pub(crate) fn validate_models_dev_provider_namespace(namespace: &str) -> anyhow::Result<String> {
262 let namespace = namespace.trim();
263 if looks_like_secret_value(namespace) {
264 anyhow::bail!(
265 "models.dev provider namespace looks like a secret value; enter a namespace such as openai"
266 );
267 }
268 if namespace.len() > 63
269 || namespace.is_empty()
270 || !namespace.as_bytes()[0].is_ascii_lowercase()
271 || namespace.ends_with('-')
272 || !namespace
273 .chars()
274 .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
275 {
276 anyhow::bail!(
277 "models.dev provider namespace must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
278 );
279 }
280 Ok(namespace.to_string())
281}
282
283pub(crate) fn derive_custom_provider_id(label: &str) -> anyhow::Result<String> {
284 let mut id = String::new();
285 let mut last_was_separator = false;
286 for ch in label.trim().chars() {
287 if ch.is_ascii_alphanumeric() {
288 id.push(ch.to_ascii_lowercase());
289 last_was_separator = false;
290 } else if !last_was_separator && !id.is_empty() {
291 id.push('-');
292 last_was_separator = true;
293 }
294 }
295 while id.ends_with('-') {
296 id.pop();
297 }
298 validate_custom_provider_id(&id)
299 .map_err(|_| anyhow::anyhow!("custom provider label must derive a provider id matching ^[a-z][a-z0-9-]{{0,62}}$ and must not be reserved"))
300}
301
302pub(crate) fn normalize_custom_provider_base_url(input: &str) -> anyhow::Result<String> {
303 let value = input.trim().trim_end_matches('/');
304 let parsed = reqwest::Url::parse(value)
305 .map_err(|_| anyhow::anyhow!("custom provider base URL must be a valid URL"))?;
306 if !parsed.username().is_empty() || parsed.password().is_some() {
307 anyhow::bail!("custom provider base URL must not include URL credentials or userinfo");
308 }
309 if parsed.query().is_some() || parsed.fragment().is_some() {
310 anyhow::bail!("custom provider base URL must not include query parameters or fragments");
311 }
312 let path = parsed.path().trim_end_matches('/');
313 if path.ends_with("/responses")
314 || path.ends_with("/models")
315 || path.ends_with("/completions")
316 || path.ends_with("/chat/completions")
317 {
318 anyhow::bail!("custom provider base URL must be an API root, not an endpoint URL");
319 }
320 match parsed.scheme() {
321 "https" | "http" => Ok(value.to_string()),
322 _ => anyhow::bail!("custom provider base URL must use http:// or https://"),
323 }
324}
325
326pub(crate) fn make_custom_provider_config(
327 label: &str,
328 base_url: &str,
329 api_key_env_var: &str,
330) -> anyhow::Result<CustomProviderConfig> {
331 let label = label.trim();
332 validate_custom_provider_label(label)?;
333 Ok(CustomProviderConfig {
334 label: label.to_string(),
335 base_url: normalize_custom_provider_base_url(base_url)?,
336 api_key_env_var: validate_optional_env_var_name(api_key_env_var)?,
337 models_dev_provider: None,
338 fast_mode: None,
339 use_responses_endpoint: default_use_responses_endpoint(),
340 supports_text_verbosity: false,
341 reasoning_protocol: CustomReasoningProtocol::default(),
342 extra_models: Vec::new(),
343 request_headers: BTreeMap::new(),
344 })
345}
346
347fn validate_custom_provider_fast_mode(fast_mode: &CustomProviderFastMode) -> anyhow::Result<()> {
348 validate_fast_service_tier(&fast_mode.service_tier)?;
349 validate_fast_models(&fast_mode.models)?;
350 Ok(())
351}
352
353fn validate_fast_service_tier(service_tier: &str) -> anyhow::Result<String> {
354 let service_tier = service_tier.trim();
355 if service_tier.is_empty()
356 || service_tier.len() > 64
357 || !service_tier.bytes().all(|byte| {
358 byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-' || byte == b'.'
359 })
360 {
361 anyhow::bail!(
362 "fast_mode.service_tier must be a non-empty ASCII identifier of at most 64 characters"
363 );
364 }
365 if looks_like_secret_value(service_tier) {
366 anyhow::bail!("fast_mode.service_tier must not look like a secret value");
367 }
368 Ok(service_tier.to_string())
369}
370
371fn validate_fast_models(models: &[String]) -> anyhow::Result<()> {
372 if models.is_empty() || models.len() > 64 {
373 anyhow::bail!("fast_mode.models must contain 1 to 64 model ids");
374 }
375 if models.iter().any(|model| model == "*") && models.len() != 1 {
376 anyhow::bail!("fast_mode.models wildcard must be the sole member");
377 }
378 let mut seen = BTreeSet::new();
379 for model in models {
380 if model.is_empty()
381 || model.chars().count() > 200
382 || model
383 .chars()
384 .any(|ch| ch.is_whitespace() || ch.is_control())
385 {
386 anyhow::bail!(
387 "fast_mode.models entries must be non-empty model ids of at most 200 characters without whitespace or control characters"
388 );
389 }
390 if model != "*" && looks_like_secret_value(model) {
391 anyhow::bail!("fast_mode.models entries must not look like secret values");
392 }
393 if !seen.insert(model.as_str()) {
394 anyhow::bail!("fast_mode.models must not contain duplicate model ids");
395 }
396 }
397 Ok(())
398}
399
400#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
401pub struct CustomProviderFastMode {
402 pub service_tier: String,
403 pub models: Vec<String>,
404}
405
406#[derive(JsonSchema)]
407#[expect(dead_code)]
408struct CustomProviderFastModeSchema {
409 #[schemars(regex(pattern = r"^\s*[A-Za-z0-9_.-]{1,64}\s*$"))]
410 service_tier: String,
411 #[schemars(
412 length(min = 1, max = 64),
413 inner(regex(pattern = r"^\S{1,200}$")),
414 transform = add_fast_models_schema_constraints
415 )]
416 models: Vec<String>,
417}
418
419fn add_fast_models_schema_constraints(schema: &mut schemars::Schema) {
420 let object = schema.ensure_object();
421 object.insert("uniqueItems".to_string(), serde_json::json!(true));
422 object.insert(
423 "oneOf".to_string(),
424 serde_json::json!([
425 {
426 "contains": {"pattern": r"^\*$"},
427 "maxItems": 1
428 },
429 {
430 "not": {"contains": {"pattern": r"^\*$"}}
431 }
432 ]),
433 );
434}
435
436impl JsonSchema for CustomProviderFastMode {
437 fn schema_name() -> std::borrow::Cow<'static, str> {
438 "CustomProviderFastMode".into()
439 }
440
441 fn json_schema(generator: &mut schemars::SchemaGenerator) -> schemars::Schema {
442 CustomProviderFastModeSchema::json_schema(generator)
443 }
444}
445
446fn is_gpt_like(value: &CustomReasoningProtocol) -> bool {
447 *value == CustomReasoningProtocol::GptLike
448}
449
450fn is_false(value: &bool) -> bool {
451 !*value
452}
453
454fn deserialize_optional_env_var<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
455where
456 D: serde::Deserializer<'de>,
457{
458 let value = Option::<String>::deserialize(deserializer)?;
459 Ok(value.and_then(|value| {
460 let trimmed = value.trim();
461 if trimmed.is_empty() {
462 None
463 } else {
464 Some(trimmed.to_string())
465 }
466 }))
467}
468
469fn deserialize_optional_fast_mode<'de, D>(
470 deserializer: D,
471) -> Result<Option<CustomProviderFastMode>, D::Error>
472where
473 D: serde::Deserializer<'de>,
474{
475 let value = Option::<CustomProviderFastMode>::deserialize(deserializer)?;
476 Ok(value.map(|mut fast| {
477 fast.service_tier = fast.service_tier.trim().to_string();
478 fast
479 }))
480}
481
482fn deserialize_optional_models_dev_provider<'de, D>(
483 deserializer: D,
484) -> Result<Option<String>, D::Error>
485where
486 D: serde::Deserializer<'de>,
487{
488 let value = Option::<String>::deserialize(deserializer)?;
489 Ok(value.and_then(|value| {
490 let trimmed = value.trim();
491 if trimmed.is_empty() {
492 None
493 } else {
494 Some(trimmed.to_string())
495 }
496 }))
497}
498
499fn deserialize_extra_models<'de, D>(deserializer: D) -> Result<Vec<String>, D::Error>
500where
501 D: serde::Deserializer<'de>,
502{
503 let values = Vec::<String>::deserialize(deserializer)?;
504 Ok(values
505 .into_iter()
506 .map(|value| value.trim().to_string())
507 .collect())
508}