Skip to main content

magi_code/config/
custom_provider_config.rs

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}