magi-code 0.63.4

Repository-aware CLI coding agent for terminal work
Documentation
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::BTreeSet;

use super::Settings;

#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
#[serde(rename_all = "kebab-case")]
pub enum CustomReasoningProtocol {
    #[default]
    GptLike,
    AnthropicLike,
}

#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct CustomProviderConfig {
    pub label: String,
    pub base_url: String,
    #[serde(
        default,
        deserialize_with = "deserialize_optional_env_var",
        skip_serializing_if = "Option::is_none"
    )]
    pub api_key_env_var: Option<String>,
    #[serde(
        default,
        deserialize_with = "deserialize_optional_models_dev_provider",
        skip_serializing_if = "Option::is_none"
    )]
    pub models_dev_provider: Option<String>,
    #[serde(default, skip_serializing_if = "is_false")]
    pub use_responses_endpoint: bool,
    #[serde(default, skip_serializing_if = "is_false")]
    pub supports_text_verbosity: bool,
    #[serde(default, skip_serializing_if = "is_gpt_like")]
    pub reasoning_protocol: CustomReasoningProtocol,
    #[serde(
        default,
        deserialize_with = "deserialize_extra_models",
        skip_serializing_if = "Vec::is_empty"
    )]
    pub extra_models: Vec<String>,
}

#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CustomProviderRegistration {
    pub id: String,
    pub config: CustomProviderConfig,
}

pub(super) fn validate_custom_provider_settings(settings: &Settings) -> anyhow::Result<()> {
    for (id, custom) in &settings.custom_providers {
        validate_custom_provider_id(id).map_err(|error| {
            anyhow::anyhow!("custom provider '{id}' has invalid provider id: {error}")
        })?;
        validate_custom_provider_label(&custom.label).map_err(|error| {
            anyhow::anyhow!("custom provider '{id}' has invalid label: {error}")
        })?;
        normalize_custom_provider_base_url(&custom.base_url).map_err(|error| {
            anyhow::anyhow!("custom provider '{id}' has invalid base_url: {error}")
        })?;
        if let Some(env_var) = &custom.api_key_env_var {
            validate_env_var_name(env_var).map_err(|error| {
                anyhow::anyhow!("custom provider '{id}' has invalid api_key_env_var: {error}")
            })?;
        }
        if let Some(models_dev_provider) = &custom.models_dev_provider {
            validate_models_dev_provider_namespace(models_dev_provider).map_err(|error| {
                anyhow::anyhow!("custom provider '{id}' has invalid models_dev_provider: {error}")
            })?;
        }
        normalized_extra_models(&custom.extra_models).map_err(|error| {
            anyhow::anyhow!("custom provider '{id}' has invalid extra_models: {error}")
        })?;
    }
    Ok(())
}

fn validate_custom_provider_label(label: &str) -> anyhow::Result<()> {
    let label = label.trim();
    if label.is_empty() || label.len() > 100 {
        anyhow::bail!("custom provider label must be non-empty and at most 100 characters");
    }
    if looks_like_secret_label(label) {
        anyhow::bail!("custom provider label must not look like a secret value");
    }
    Ok(())
}

fn looks_like_secret_label(value: &str) -> bool {
    let value = value.trim();
    value.starts_with("sk-")
        || value.starts_with("Bearer ")
        || value.contains('=')
        || (value.len() >= 48
            && value
                .chars()
                .filter(|ch| ch.is_ascii_alphanumeric())
                .count()
                >= 40)
}

pub fn validate_custom_provider_id(id: &str) -> anyhow::Result<String> {
    let id = id.trim();
    if matches!(
        id,
        crate::providers::OPENAI_CODEX_PROVIDER
            | crate::providers::ANTHROPIC_PROVIDER
            | crate::providers::CLAUDE_CODE_PROVIDER
            | "openai"
    ) {
        anyhow::bail!("custom provider id '{id}' is reserved");
    }
    if id.len() > 63
        || id.is_empty()
        || !id.as_bytes()[0].is_ascii_lowercase()
        || id.ends_with('-')
        || !id
            .chars()
            .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
    {
        anyhow::bail!(
            "custom provider id must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
        );
    }
    Ok(id.to_string())
}

pub fn looks_like_secret_value(value: &str) -> bool {
    let value = value.trim();
    value.starts_with("sk-")
        || value.starts_with("Bearer ")
        || value.contains('=')
        || value.chars().any(char::is_whitespace)
        || (value.len() >= 48
            && value
                .chars()
                .filter(|ch| ch.is_ascii_alphanumeric())
                .count()
                >= 40)
}

pub fn validate_env_var_name(name: &str) -> anyhow::Result<String> {
    let name = name.trim();
    if looks_like_secret_value(name) {
        anyhow::bail!(
            "API key environment variable name looks like a secret value; enter a variable name such as CUSTOM_PROVIDER_API_KEY"
        );
    }
    if name.is_empty()
        || !(name.as_bytes()[0].is_ascii_uppercase() || name.as_bytes()[0] == b'_')
        || !name
            .chars()
            .all(|ch| ch.is_ascii_uppercase() || ch.is_ascii_digit() || ch == '_')
    {
        anyhow::bail!("API key environment variable name must match ^[A-Z_][A-Z0-9_]*$");
    }
    Ok(name.to_string())
}

pub fn validate_optional_env_var_name(name: &str) -> anyhow::Result<Option<String>> {
    if name.trim().is_empty() {
        return Ok(None);
    }
    validate_env_var_name(name).map(Some)
}

pub(crate) fn normalized_extra_models(extra_models: &[String]) -> anyhow::Result<Vec<String>> {
    if extra_models.len() > 64 {
        anyhow::bail!("extra_models must contain at most 64 model ids");
    }
    let mut seen = BTreeSet::new();
    let mut normalized = Vec::new();
    for model in extra_models {
        let model = model.trim();
        if model.is_empty() {
            anyhow::bail!("extra_models entries must be non-empty");
        }
        if model.len() > 200 {
            anyhow::bail!("extra_models entries must be at most 200 bytes");
        }
        if model
            .chars()
            .any(|ch| ch.is_ascii_control() || ch.is_ascii_whitespace())
        {
            anyhow::bail!(
                "extra_models entries must not contain ASCII control characters or whitespace"
            );
        }
        if looks_like_secret_value(model) {
            anyhow::bail!("extra_models entries must not look like secret values");
        }
        if seen.insert(model.to_string()) {
            normalized.push(model.to_string());
        }
    }
    Ok(normalized)
}

pub fn validate_models_dev_provider_namespace(namespace: &str) -> anyhow::Result<String> {
    let namespace = namespace.trim();
    if looks_like_secret_value(namespace) {
        anyhow::bail!(
            "models.dev provider namespace looks like a secret value; enter a namespace such as openai"
        );
    }
    if namespace.len() > 63
        || namespace.is_empty()
        || !namespace.as_bytes()[0].is_ascii_lowercase()
        || namespace.ends_with('-')
        || !namespace
            .chars()
            .all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
    {
        anyhow::bail!(
            "models.dev provider namespace must match ^[a-z][a-z0-9-]{{0,62}}$ with no trailing hyphen"
        );
    }
    Ok(namespace.to_string())
}

pub fn derive_custom_provider_id(label: &str) -> anyhow::Result<String> {
    let mut id = String::new();
    let mut last_was_separator = false;
    for ch in label.trim().chars() {
        if ch.is_ascii_alphanumeric() {
            id.push(ch.to_ascii_lowercase());
            last_was_separator = false;
        } else if !last_was_separator && !id.is_empty() {
            id.push('-');
            last_was_separator = true;
        }
    }
    while id.ends_with('-') {
        id.pop();
    }
    validate_custom_provider_id(&id)
        .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"))
}

pub fn normalize_custom_provider_base_url(input: &str) -> anyhow::Result<String> {
    let value = input.trim().trim_end_matches('/');
    let parsed = reqwest::Url::parse(value)
        .map_err(|_| anyhow::anyhow!("custom provider base URL must be a valid URL"))?;
    if !parsed.username().is_empty() || parsed.password().is_some() {
        anyhow::bail!("custom provider base URL must not include URL credentials or userinfo");
    }
    if parsed.query().is_some() || parsed.fragment().is_some() {
        anyhow::bail!("custom provider base URL must not include query parameters or fragments");
    }
    let path = parsed.path().trim_end_matches('/');
    if path.ends_with("/responses")
        || path.ends_with("/models")
        || path.ends_with("/completions")
        || path.ends_with("/chat/completions")
    {
        anyhow::bail!("custom provider base URL must be an API root, not an endpoint URL");
    }
    match parsed.scheme() {
        "https" | "http" => Ok(value.to_string()),
        _ => anyhow::bail!("custom provider base URL must use http:// or https://"),
    }
}

pub fn make_custom_provider_config(
    label: &str,
    base_url: &str,
    api_key_env_var: &str,
) -> anyhow::Result<CustomProviderConfig> {
    let label = label.trim();
    validate_custom_provider_label(label)?;
    Ok(CustomProviderConfig {
        label: label.to_string(),
        base_url: normalize_custom_provider_base_url(base_url)?,
        api_key_env_var: validate_optional_env_var_name(api_key_env_var)?,
        models_dev_provider: None,
        use_responses_endpoint: false,
        supports_text_verbosity: false,
        reasoning_protocol: CustomReasoningProtocol::default(),
        extra_models: Vec::new(),
    })
}

fn is_gpt_like(value: &CustomReasoningProtocol) -> bool {
    *value == CustomReasoningProtocol::GptLike
}

fn is_false(value: &bool) -> bool {
    !*value
}

fn deserialize_optional_env_var<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
where
    D: serde::Deserializer<'de>,
{
    let value = Option::<String>::deserialize(deserializer)?;
    Ok(value.and_then(|value| {
        let trimmed = value.trim();
        if trimmed.is_empty() {
            None
        } else {
            Some(trimmed.to_string())
        }
    }))
}

fn deserialize_optional_models_dev_provider<'de, D>(
    deserializer: D,
) -> Result<Option<String>, D::Error>
where
    D: serde::Deserializer<'de>,
{
    let value = Option::<String>::deserialize(deserializer)?;
    Ok(value.and_then(|value| {
        let trimmed = value.trim();
        if trimmed.is_empty() {
            None
        } else {
            Some(trimmed.to_string())
        }
    }))
}

fn deserialize_extra_models<'de, D>(deserializer: D) -> Result<Vec<String>, D::Error>
where
    D: serde::Deserializer<'de>,
{
    let values = Vec::<String>::deserialize(deserializer)?;
    Ok(values
        .into_iter()
        .map(|value| value.trim().to_string())
        .collect())
}