Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-ai/src/models.rs

Raw
//! Offline model catalog, capabilities, and deterministic resolution
//! (`SMH-SPEC-SPEC0001`, Providers).
//!
//! The reviewed catalog ships as checked-in JSON. Resolution is pure data
//! navigation with no network: aliases recurse, groups expand to ordered
//! failover plans, and cycles or dangling references fail at load with the
//! field path that caused them.

use serde::{Deserialize, Serialize};
use smith::error::{ProviderFault, Result, SmithError};
use smith::provider::{ProviderParams, ProviderRequest};
use std::collections::HashMap;

/// Which adapter family serves a model.
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug, Serialize, Deserialize)]
pub enum ProviderKind {
    /// OpenAI-compatible chat completions.
    OpenAiCompatible,
    /// Anthropic Messages.
    Anthropic,
    /// Google Gemini.
    Google,
    /// Plugin-provided.
    Plugin,
}

impl ProviderKind {
    /// Stable lowercase file name for this provider.
    #[must_use]
    pub const fn slug(self) -> &'static str {
        match self {
            Self::OpenAiCompatible => "openai",
            Self::Anthropic => "anthropic",
            Self::Google => "google",
            Self::Plugin => "plugin",
        }
    }

    /// Parse a provider from its slug.
    #[must_use]
    pub fn from_slug(slug: &str) -> Option<Self> {
        match slug {
            "openai" => Some(Self::OpenAiCompatible),
            "anthropic" => Some(Self::Anthropic),
            "google" => Some(Self::Google),
            "plugin" => Some(Self::Plugin),
            _ => None,
        }
    }
}

/// What a model explicitly supports, used to validate requests.
#[expect(
    clippy::struct_excessive_bools,
    reason = "capabilities are independent advertised switches"
)]
#[derive(Clone, PartialEq, Eq, Debug, Serialize, Deserialize)]
pub struct ModelCapabilities {
    /// Context window in tokens.
    pub context_window: u64,
    /// Output limit in tokens.
    pub max_output_tokens: u64,
    /// Accepts image inputs.
    pub input_images: bool,
    /// Accepts tool definitions.
    pub tool_use: bool,
    /// Supports thinking output.
    pub thinking: bool,
    /// Supports strict or grammar-constrained tool input.
    pub strict_tool_input: bool,
    /// Supports provider-side prompt caching.
    pub cache_control: bool,
}

/// One catalog model entry.
#[derive(Clone, PartialEq, Eq, Debug, Serialize, Deserialize)]
pub struct ModelEntry {
    /// Model identifier used in requests.
    pub id: String,
    /// Adapter family.
    pub provider: ProviderKind,
    /// Advertised capabilities.
    pub capabilities: ModelCapabilities,
}

/// A reviewed offline catalog plus its resolution tables.
#[derive(Clone, PartialEq, Eq, Debug, Serialize, Deserialize)]
pub struct Catalog {
    /// Model entries by id.
    pub models: HashMap<String, ModelEntry>,
    /// Alias to model id or group name.
    pub aliases: HashMap<String, String>,
    /// Group name to ordered member model ids.
    pub groups: HashMap<String, Vec<String>>,
}

/// The reviewed offline catalog, checked in as data.
const BUILTIN: &str = include_str!("catalog.json");

impl Catalog {
    /// Load the built-in reviewed catalog, validating all references.
    ///
    /// # Errors
    ///
    /// Returns a [`SmithError::Config`] naming the offending field path when
    /// the data references a missing model, alias, or group, or contains a
    /// direct alias cycle.
    pub fn builtin() -> Result<Self> {
        Self::from_json(BUILTIN)
    }

    /// Parse and validate catalog JSON.
    ///
    /// # Errors
    ///
    /// Same failures as [`Catalog::builtin`].
    pub fn from_json(json: &str) -> Result<Self> {
        let catalog: Self =
            serde_json::from_str(json).map_err(|e| invalid("catalog", &e.to_string()))?;
        // Deterministic diagnostics: validation walks aliases in sorted order.
        let mut alias_names: Vec<&String> = catalog.aliases.keys().collect();
        alias_names.sort();
        for alias in &alias_names {
            let target = &catalog.aliases[*alias];
            if !catalog.models.contains_key(target)
                && !catalog.groups.contains_key(target)
                && !catalog.aliases.contains_key(target)
            {
                return Err(invalid(
                    &format!("catalog.aliases.{alias}"),
                    &format!("target '{target}' is neither a model nor a group"),
                ));
            }
        }
        for (group, members) in &catalog.groups {
            for (position, member) in members.iter().enumerate() {
                if !catalog.models.contains_key(member) {
                    return Err(invalid(
                        &format!("catalog.groups.{group}[{position}]"),
                        &format!("member '{member}' is not a model"),
                    ));
                }
            }
        }
        for alias in &alias_names {
            catalog.resolve_alias(alias)?;
        }
        Ok(catalog)
    }

    fn resolve_alias(&self, alias: &str) -> Result<String> {
        let mut chain = vec![alias.to_string()];
        let mut current = alias;
        while let Some(next) = self.aliases.get(current) {
            if chain.iter().any(|step| step == next) {
                chain.push(next.clone());
                return Err(invalid(
                    &format!("catalog.aliases.{alias}"),
                    &format!("alias cycle: {}", chain.join(" -> ")),
                ));
            }
            chain.push(next.clone());
            current = next;
        }
        Ok(current.to_string())
    }

    /// Resolve a request target into an ordered failover plan.
    ///
    /// A model id resolves to itself; an alias resolves to its referent; a
    /// group resolves to its members in declared order. Resolution performs
    /// no network access.
    ///
    /// # Errors
    ///
    /// Returns a [`SmithError::Config`] naming the target when it is unknown
    /// or cyclic.
    pub fn resolve_plan(&self, requested: &str) -> Result<Vec<&ModelEntry>> {
        let resolved = self.resolve_alias(requested)?;
        let ids: Vec<&String> = self
            .groups
            .get(&resolved)
            .map_or_else(|| vec![&resolved], |members| members.iter().collect());
        ids.into_iter()
            .map(|id| {
                self.models.get(id).ok_or_else(|| {
                    invalid(
                        &format!("catalog.resolve.{requested}"),
                        &format!("'{id}' is not a model"),
                    )
                })
            })
            .collect()
    }

    /// Validate that a request uses only capabilities the model advertises.
    ///
    /// # Errors
    ///
    /// Returns a [`SmithError::Provider`] with fault
    /// [`ProviderFault::Invalid`] naming the offending request field when the
    /// model does not advertise what the request asks for.
    pub fn validate_request(&self, model: &str, request: &ProviderRequest) -> Result<()> {
        let entry =
            self.resolve_plan(model)?.first().copied().ok_or_else(|| {
                invalid(&format!("catalog.resolve.{model}"), "empty resolution plan")
            })?;
        let caps = &entry.capabilities;
        if !request.tools.is_empty() && !caps.tool_use {
            return Err(invalid_request(
                "tools",
                &format!("{} does not advertise tool use", entry.id),
            ));
        }
        let ProviderParams {
            max_output_tokens,
            temperature,
            thinking,
        } = request.params;
        if let Some(requested) = max_output_tokens
            && requested > caps.max_output_tokens
        {
            return Err(invalid_request(
                "params.max_output_tokens",
                &format!(
                    "{requested} exceeds the advertised limit of {}",
                    caps.max_output_tokens
                ),
            ));
        }
        if let Some(requested) = temperature
            && requested > 2
        {
            return Err(invalid_request(
                "params.temperature",
                "temperature above 2 is not a valid sampling value",
            ));
        }
        if thinking && !caps.thinking {
            return Err(invalid_request(
                "params.thinking",
                &format!("{} does not advertise thinking", entry.id),
            ));
        }
        Ok(())
    }
}

fn invalid(field: &str, message: &str) -> SmithError {
    SmithError::Config {
        field: field.to_string(),
        message: message.to_string(),
    }
}

fn invalid_request(field: &str, message: &str) -> SmithError {
    SmithError::Provider {
        fault: ProviderFault::Invalid {
            field: Some(field.to_string()),
            message: message.to_string(),
        },
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use smith::message::{Message, Role};
    use smith::provider::ProviderParams;
    use smith::tool::ToolMetadata;

    const TEST_CATALOG: &str = r#"{
        "models": {
            "mock-smith": {
                "id": "mock-smith",
                "provider": "OpenAiCompatible",
                "capabilities": {
                    "context_window": 8192,
                    "max_output_tokens": 1024,
                    "input_images": false,
                    "tool_use": true,
                    "thinking": false,
                    "strict_tool_input": false,
                    "cache_control": false
                }
            },
            "thinker": {
                "id": "thinker",
                "provider": "Anthropic",
                "capabilities": {
                    "context_window": 200000,
                    "max_output_tokens": 8192,
                    "input_images": true,
                    "tool_use": true,
                    "thinking": true,
                    "strict_tool_input": false,
                    "cache_control": true
                }
            },
            "no-tools": {
                "id": "no-tools",
                "provider": "Google",
                "capabilities": {
                    "context_window": 8192,
                    "max_output_tokens": 1024,
                    "input_images": false,
                    "tool_use": false,
                    "thinking": false,
                    "strict_tool_input": false,
                    "cache_control": false
                }
            }
        },
        "aliases": { "smart": "thinker" },
        "groups": { "default": ["thinker", "mock-smith"] }
    }"#;

    #[expect(
        clippy::unwrap_used,
        reason = "catalog fixtures are valid by construction"
    )]
    fn catalog() -> Catalog {
        Catalog::from_json(TEST_CATALOG).unwrap()
    }

    #[test]
    fn builtin_catalog_loads_and_resolves() {
        let catalog = Catalog::builtin().unwrap();
        let _plan = catalog.resolve_plan("default");
        // Every builtin id resolves to itself.
        for id in catalog.models.keys() {
            assert!(catalog.resolve_plan(id).is_ok(), "id {id}");
        }
    }

    #[test]
    fn alias_resolves_to_model_and_group_expands_in_order() {
        let catalog = catalog();
        let plan = catalog.resolve_plan("smart").unwrap();
        assert_eq!(plan.len(), 1);
        assert_eq!(plan[0].id, "thinker");
        let plan = catalog.resolve_plan("default").unwrap();
        assert_eq!(
            plan.iter()
                .map(|entry| entry.id.as_str())
                .collect::<Vec<_>>(),
            vec!["thinker", "mock-smith"]
        );
    }

    #[test]
    fn unknown_targets_and_cycles_fail_with_field_paths() {
        let catalog = catalog();
        let err = catalog.resolve_plan("nope").unwrap_err();
        assert_eq!(err.code(), "CONFIG");
        let cyclic = r#"{
            "models": {},
            "aliases": { "a": "b", "b": "a" },
            "groups": {}
        }"#;
        let err = Catalog::from_json(cyclic).unwrap_err();
        assert_eq!(
            err.to_string(),
            "config [catalog.aliases.a]: alias cycle: a -> b -> a"
        );
        let dangling = r#"{
            "models": {},
            "aliases": { "smart": "ghost" },
            "groups": {}
        }"#;
        let err = Catalog::from_json(dangling).unwrap_err();
        assert!(err.to_string().contains("catalog.aliases.smart"));
    }

    #[test]
    fn requests_cannot_use_unadvertised_capabilities() {
        let catalog = catalog();
        let tools = vec![ToolMetadata::new("read", "Read a file")];
        let mut request =
            ProviderRequest::new("mock-smith", vec![Message::with_text(Role::User, "hi")])
                .with_tools(tools);
        // mock-smith advertises tools: allowed.
        assert!(catalog.validate_request("mock-smith", &request).is_ok());
        let tools = vec![ToolMetadata::new("read", "Read a file")];
        let _ = tools;
        // Thinking not advertised: rejected with the field path.
        request.params = ProviderParams {
            thinking: true,
            ..ProviderParams::default()
        };
        let err = catalog
            .validate_request("mock-smith", &request)
            .unwrap_err();
        assert_eq!(err.code(), "PROVIDER_INVALID");
        assert!(err.to_string().contains("params.thinking"));
        // Output limit exceeded: rejected.
        request.params = ProviderParams {
            max_output_tokens: Some(999_999),
            thinking: false,
            ..ProviderParams::default()
        };
        assert!(catalog.validate_request("mock-smith", &request).is_err());
        // thinker advertises thinking but not an unlimited budget: also rejected.
        assert!(catalog.validate_request("thinker", &request).is_err());
        request.params = ProviderParams {
            max_output_tokens: Some(4096),
            thinking: true,
            ..ProviderParams::default()
        };
        assert!(catalog.validate_request("thinker", &request).is_ok());
    }

    #[test]
    fn provider_slugs_round_trip() {
        for (kind, slug) in [
            (ProviderKind::OpenAiCompatible, "openai"),
            (ProviderKind::Anthropic, "anthropic"),
            (ProviderKind::Google, "google"),
            (ProviderKind::Plugin, "plugin"),
        ] {
            assert_eq!(kind.slug(), slug);
            assert_eq!(ProviderKind::from_slug(slug), Some(kind));
        }
        assert_eq!(ProviderKind::from_slug("xyzzy"), None);
    }

    #[test]
    fn validation_holds_exactly_at_each_limit() {
        let catalog = catalog();
        let params = |max_output_tokens, temperature| ProviderParams {
            max_output_tokens,
            temperature,
            thinking: false,
        };
        let none = ProviderParams::default();
        // (model, with tools, params, rejected field)
        let cases = [
            ("mock-smith", false, params(Some(1024), None), None),
            (
                "mock-smith",
                false,
                params(Some(1025), None),
                Some("params.max_output_tokens"),
            ),
            ("mock-smith", false, params(None, Some(0)), None),
            ("mock-smith", false, params(None, Some(2)), None),
            (
                "mock-smith",
                false,
                params(None, Some(3)),
                Some("params.temperature"),
            ),
            ("no-tools", false, none, None),
            ("no-tools", true, none, Some("tools")),
        ];
        for (model, with_tools, params, rejected) in cases {
            let mut request =
                ProviderRequest::new(model, vec![Message::with_text(Role::User, "hi")]);
            if with_tools {
                request = request.with_tools(vec![ToolMetadata::new("read", "Read a file")]);
            }
            request.params = params;
            let field = match catalog.validate_request(model, &request) {
                Ok(()) => None,
                Err(SmithError::Provider {
                    fault: ProviderFault::Invalid { field, .. },
                }) => field,
                Err(other) => Some(other.to_string()),
            };
            assert_eq!(
                field.as_deref(),
                rejected,
                "{model} tools={with_tools} {params:?}"
            );
        }
    }
}