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 | "claude-code"
171 ) {
172 anyhow::bail!("custom provider id '{id}' is reserved");
173 }
174 if id.len() > 63
175 || id.is_empty()
176 || !id.as_bytes()[0].is_ascii_lowercase()
177 || id.ends_with('-')
178 || !id
179 .chars()
180 .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
181 {
182 anyhow::bail!(
183 "custom provider id must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
184 );
185 }
186 Ok(id.to_string())
187}
188
189pub(crate) fn looks_like_secret_value(value: &str) -> bool {
190 let value = value.trim();
191 value.starts_with("sk-")
192 || value.starts_with("Bearer ")
193 || value.contains('=')
194 || value.chars().any(char::is_whitespace)
195 || (value.len() >= 48
196 && value
197 .chars()
198 .filter(|ch| ch.is_ascii_alphanumeric())
199 .count()
200 >= 40)
201}
202
203pub(crate) fn validate_env_var_name(name: &str) -> anyhow::Result<String> {
204 let name = name.trim();
205 if looks_like_secret_value(name) {
206 anyhow::bail!(
207 "API key environment variable name looks like a secret value; enter a variable name such as CUSTOM_PROVIDER_API_KEY"
208 );
209 }
210 if name.is_empty()
211 || !(name.as_bytes()[0].is_ascii_uppercase() || name.as_bytes()[0] == b'_')
212 || !name
213 .chars()
214 .all(|ch| ch.is_ascii_uppercase() || ch.is_ascii_digit() || ch == '_')
215 {
216 anyhow::bail!("API key environment variable name must match ^[A-Z_][A-Z0-9_]*$");
217 }
218 Ok(name.to_string())
219}
220
221pub(crate) fn validate_optional_env_var_name(name: &str) -> anyhow::Result<Option<String>> {
222 if name.trim().is_empty() {
223 return Ok(None);
224 }
225 validate_env_var_name(name).map(Some)
226}
227
228pub(crate) fn normalized_extra_models(extra_models: &[String]) -> anyhow::Result<Vec<String>> {
229 if extra_models.len() > 64 {
230 anyhow::bail!("extra_models must contain at most 64 model ids");
231 }
232 let mut seen = BTreeSet::new();
233 let mut normalized = Vec::new();
234 for model in extra_models {
235 let model = model.trim();
236 if model.is_empty() {
237 anyhow::bail!("extra_models entries must be non-empty");
238 }
239 if model.len() > 200 {
240 anyhow::bail!("extra_models entries must be at most 200 bytes");
241 }
242 if model
243 .chars()
244 .any(|ch| ch.is_ascii_control() || ch.is_ascii_whitespace())
245 {
246 anyhow::bail!(
247 "extra_models entries must not contain ASCII control characters or whitespace"
248 );
249 }
250 if looks_like_secret_value(model) {
251 anyhow::bail!("extra_models entries must not look like secret values");
252 }
253 if seen.insert(model.to_string()) {
254 normalized.push(model.to_string());
255 }
256 }
257 Ok(normalized)
258}
259
260pub(crate) fn validate_models_dev_provider_namespace(namespace: &str) -> anyhow::Result<String> {
261 let namespace = namespace.trim();
262 if looks_like_secret_value(namespace) {
263 anyhow::bail!(
264 "models.dev provider namespace looks like a secret value; enter a namespace such as openai"
265 );
266 }
267 if namespace.len() > 63
268 || namespace.is_empty()
269 || !namespace.as_bytes()[0].is_ascii_lowercase()
270 || namespace.ends_with('-')
271 || !namespace
272 .chars()
273 .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
274 {
275 anyhow::bail!(
276 "models.dev provider namespace must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
277 );
278 }
279 Ok(namespace.to_string())
280}
281
282pub(crate) fn derive_custom_provider_id(label: &str) -> anyhow::Result<String> {
283 let mut id = String::new();
284 let mut last_was_separator = false;
285 for ch in label.trim().chars() {
286 if ch.is_ascii_alphanumeric() {
287 id.push(ch.to_ascii_lowercase());
288 last_was_separator = false;
289 } else if !last_was_separator && !id.is_empty() {
290 id.push('-');
291 last_was_separator = true;
292 }
293 }
294 while id.ends_with('-') {
295 id.pop();
296 }
297 validate_custom_provider_id(&id)
298 .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"))
299}
300
301pub(crate) fn normalize_custom_provider_base_url(input: &str) -> anyhow::Result<String> {
302 let value = input.trim().trim_end_matches('/');
303 let parsed = reqwest::Url::parse(value)
304 .map_err(|_| anyhow::anyhow!("custom provider base URL must be a valid URL"))?;
305 if !parsed.username().is_empty() || parsed.password().is_some() {
306 anyhow::bail!("custom provider base URL must not include URL credentials or userinfo");
307 }
308 if parsed.query().is_some() || parsed.fragment().is_some() {
309 anyhow::bail!("custom provider base URL must not include query parameters or fragments");
310 }
311 let path = parsed.path().trim_end_matches('/');
312 if path.ends_with("/responses")
313 || path.ends_with("/models")
314 || path.ends_with("/completions")
315 || path.ends_with("/chat/completions")
316 {
317 anyhow::bail!("custom provider base URL must be an API root, not an endpoint URL");
318 }
319 match parsed.scheme() {
320 "https" | "http" => Ok(value.to_string()),
321 _ => anyhow::bail!("custom provider base URL must use http:// or https://"),
322 }
323}
324
325pub(crate) fn make_custom_provider_config(
326 label: &str,
327 base_url: &str,
328 api_key_env_var: &str,
329) -> anyhow::Result<CustomProviderConfig> {
330 let label = label.trim();
331 validate_custom_provider_label(label)?;
332 Ok(CustomProviderConfig {
333 label: label.to_string(),
334 base_url: normalize_custom_provider_base_url(base_url)?,
335 api_key_env_var: validate_optional_env_var_name(api_key_env_var)?,
336 models_dev_provider: None,
337 fast_mode: None,
338 use_responses_endpoint: default_use_responses_endpoint(),
339 supports_text_verbosity: false,
340 reasoning_protocol: CustomReasoningProtocol::default(),
341 extra_models: Vec::new(),
342 request_headers: BTreeMap::new(),
343 })
344}
345
346fn validate_custom_provider_fast_mode(fast_mode: &CustomProviderFastMode) -> anyhow::Result<()> {
347 validate_fast_service_tier(&fast_mode.service_tier)?;
348 validate_fast_models(&fast_mode.models)?;
349 Ok(())
350}
351
352fn validate_fast_service_tier(service_tier: &str) -> anyhow::Result<String> {
353 let service_tier = service_tier.trim();
354 if service_tier.is_empty()
355 || service_tier.len() > 64
356 || !service_tier.bytes().all(|byte| {
357 byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-' || byte == b'.'
358 })
359 {
360 anyhow::bail!(
361 "fast_mode.service_tier must be a non-empty ASCII identifier of at most 64 characters"
362 );
363 }
364 if looks_like_secret_value(service_tier) {
365 anyhow::bail!("fast_mode.service_tier must not look like a secret value");
366 }
367 Ok(service_tier.to_string())
368}
369
370fn validate_fast_models(models: &[String]) -> anyhow::Result<()> {
371 if models.is_empty() || models.len() > 64 {
372 anyhow::bail!("fast_mode.models must contain 1 to 64 model ids");
373 }
374 if models.iter().any(|model| model == "*") && models.len() != 1 {
375 anyhow::bail!("fast_mode.models wildcard must be the sole member");
376 }
377 let mut seen = BTreeSet::new();
378 for model in models {
379 if model.is_empty()
380 || model.chars().count() > 200
381 || model
382 .chars()
383 .any(|ch| ch.is_whitespace() || ch.is_control())
384 {
385 anyhow::bail!(
386 "fast_mode.models entries must be non-empty model ids of at most 200 characters without whitespace or control characters"
387 );
388 }
389 if model != "*" && looks_like_secret_value(model) {
390 anyhow::bail!("fast_mode.models entries must not look like secret values");
391 }
392 if !seen.insert(model.as_str()) {
393 anyhow::bail!("fast_mode.models must not contain duplicate model ids");
394 }
395 }
396 Ok(())
397}
398
399#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
400pub struct CustomProviderFastMode {
401 pub service_tier: String,
402 pub models: Vec<String>,
403}
404
405#[derive(JsonSchema)]
406#[allow(dead_code)]
407struct CustomProviderFastModeSchema {
408 #[schemars(regex(pattern = r"^\s*[A-Za-z0-9_.-]{1,64}\s*$"))]
409 service_tier: String,
410 #[schemars(
411 length(min = 1, max = 64),
412 inner(regex(pattern = r"^\S{1,200}$")),
413 transform = add_fast_models_schema_constraints
414 )]
415 models: Vec<String>,
416}
417
418fn add_fast_models_schema_constraints(schema: &mut schemars::Schema) {
419 let object = schema.ensure_object();
420 object.insert("uniqueItems".to_string(), serde_json::json!(true));
421 object.insert(
422 "oneOf".to_string(),
423 serde_json::json!([
424 {
425 "contains": {"pattern": r"^\*$"},
426 "maxItems": 1
427 },
428 {
429 "not": {"contains": {"pattern": r"^\*$"}}
430 }
431 ]),
432 );
433}
434
435impl JsonSchema for CustomProviderFastMode {
436 fn schema_name() -> std::borrow::Cow<'static, str> {
437 "CustomProviderFastMode".into()
438 }
439
440 fn json_schema(generator: &mut schemars::SchemaGenerator) -> schemars::Schema {
441 CustomProviderFastModeSchema::json_schema(generator)
442 }
443}
444
445#[cfg(test)]
446impl CustomProviderFastMode {
447 pub(crate) fn supports_model(&self, model: &str) -> bool {
448 self.models
449 .iter()
450 .any(|candidate| candidate == "*" || candidate == model)
451 }
452}
453fn is_gpt_like(value: &CustomReasoningProtocol) -> bool {
454 *value == CustomReasoningProtocol::GptLike
455}
456
457fn is_false(value: &bool) -> bool {
458 !*value
459}
460
461fn deserialize_optional_env_var<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
462where
463 D: serde::Deserializer<'de>,
464{
465 let value = Option::<String>::deserialize(deserializer)?;
466 Ok(value.and_then(|value| {
467 let trimmed = value.trim();
468 if trimmed.is_empty() {
469 None
470 } else {
471 Some(trimmed.to_string())
472 }
473 }))
474}
475
476fn deserialize_optional_fast_mode<'de, D>(
477 deserializer: D,
478) -> Result<Option<CustomProviderFastMode>, D::Error>
479where
480 D: serde::Deserializer<'de>,
481{
482 let value = Option::<CustomProviderFastMode>::deserialize(deserializer)?;
483 Ok(value.map(|mut fast| {
484 fast.service_tier = fast.service_tier.trim().to_string();
485 fast
486 }))
487}
488
489fn deserialize_optional_models_dev_provider<'de, D>(
490 deserializer: D,
491) -> Result<Option<String>, D::Error>
492where
493 D: serde::Deserializer<'de>,
494{
495 let value = Option::<String>::deserialize(deserializer)?;
496 Ok(value.and_then(|value| {
497 let trimmed = value.trim();
498 if trimmed.is_empty() {
499 None
500 } else {
501 Some(trimmed.to_string())
502 }
503 }))
504}
505
506fn deserialize_extra_models<'de, D>(deserializer: D) -> Result<Vec<String>, D::Error>
507where
508 D: serde::Deserializer<'de>,
509{
510 let values = Vec::<String>::deserialize(deserializer)?;
511 Ok(values
512 .into_iter()
513 .map(|value| value.trim().to_string())
514 .collect())
515}
516
517#[cfg(test)]
518mod tests {
519 use super::*;
520
521 fn settings_with_fast_models(models: Vec<String>) -> Settings {
522 Settings {
523 custom_providers: [(
524 "provider".to_string(),
525 CustomProviderConfig {
526 label: "Provider".to_string(),
527 base_url: "https://example.test".to_string(),
528 api_key_env_var: None,
529 models_dev_provider: None,
530 fast_mode: Some(CustomProviderFastMode {
531 service_tier: "priority".to_string(),
532 models,
533 }),
534 use_responses_endpoint: false,
535 supports_text_verbosity: false,
536 reasoning_protocol: CustomReasoningProtocol::default(),
537 extra_models: Vec::new(),
538 request_headers: BTreeMap::new(),
539 },
540 )]
541 .into(),
542 ..Settings::default()
543 }
544 }
545
546 #[test]
547 fn fast_mode_trims_service_tier_but_preserves_exact_model_ids() {
548 let config: CustomProviderConfig = serde_json::from_value(serde_json::json!({
549 "label": "Provider", "base_url": "https://example.test",
550 "fast_mode": {"service_tier": " priority ", "models": [" model-a "]}
551 }))
552 .unwrap();
553 let fast_mode = config.fast_mode.as_ref().unwrap();
554 assert_eq!(fast_mode.service_tier, "priority");
555 assert_eq!(fast_mode.models, vec![" model-a ".to_string()]);
556 assert!(!fast_mode.supports_model("model-a"));
557 assert!(fast_mode.supports_model(" model-a "));
558 assert!(
559 validate_custom_provider_settings(&settings_with_fast_models(fast_mode.models.clone()))
560 .is_err()
561 );
562 }
563
564 #[test]
565 fn request_headers_parse_conversation_source_and_reject_owned_or_duplicate_names() {
566 let config: CustomProviderConfig = serde_json::from_value(serde_json::json!({
567 "label": "Provider",
568 "base_url": "https://example.test",
569 "request_headers": {
570 "x-opencode-session": {"source": "conversation_id"}
571 }
572 }))
573 .unwrap();
574 assert_eq!(
575 config.request_headers.get("x-opencode-session"),
576 Some(&CustomProviderHeaderValue::ConversationId)
577 );
578
579 for headers in [
580 BTreeMap::from([(
581 "Authorization".to_string(),
582 CustomProviderHeaderValue::ConversationId,
583 )]),
584 BTreeMap::from([
585 (
586 "X-Session".to_string(),
587 CustomProviderHeaderValue::ConversationId,
588 ),
589 (
590 "x-session".to_string(),
591 CustomProviderHeaderValue::ConversationId,
592 ),
593 ]),
594 ] {
595 assert!(validate_custom_provider_request_headers("provider", &headers).is_err());
596 }
597 }
598
599 #[test]
600 fn fast_mode_runtime_validation_uses_unicode_character_limits_and_rejects_whitespace() {
601 for models in [
602 vec![" model".to_string()],
603 vec!["model ".to_string()],
604 vec!["model id".to_string()],
605 vec!["model\u{2003}id".to_string()],
606 vec!["model\u{0000}id".to_string()],
607 ] {
608 assert!(validate_custom_provider_settings(&settings_with_fast_models(models)).is_err());
609 }
610 assert!(
611 validate_custom_provider_settings(&settings_with_fast_models(vec!["😀".repeat(200),]))
612 .is_ok()
613 );
614 assert!(
615 validate_custom_provider_settings(&settings_with_fast_models(vec!["😀".repeat(201),]))
616 .is_err()
617 );
618 }
619
620 #[test]
621 fn fast_mode_rejects_duplicate_models_and_non_sole_wildcard() {
622 for models in [
623 vec!["model".to_string(), "model".to_string()],
624 vec!["*".to_string(), "model".to_string()],
625 ] {
626 assert!(validate_custom_provider_settings(&settings_with_fast_models(models)).is_err());
627 }
628 }
629}