Skip to main content

specado_core/adapter/
mod.rs

1use crate::types::ProviderApi;
2use crate::types::ProviderSpec;
3
4#[derive(Debug, Clone, PartialEq, Eq)]
5pub struct AdapterSelection {
6    kind: ProviderApi,
7    match_rule: AdapterMatchRule,
8    overlays: Vec<String>,
9}
10
11impl AdapterSelection {
12    fn new(kind: ProviderApi, match_rule: AdapterMatchRule) -> Self {
13        Self {
14            kind,
15            match_rule,
16            overlays: Vec::new(),
17        }
18    }
19
20    pub fn kind(&self) -> ProviderApi {
21        self.kind
22    }
23
24    pub fn match_rule(&self) -> &AdapterMatchRule {
25        &self.match_rule
26    }
27
28    pub fn overlays(&self) -> &[String] {
29        &self.overlays
30    }
31}
32
33#[derive(Debug, Clone, PartialEq, Eq)]
34pub enum AdapterMatchRule {
35    InterfaceHint(String),
36    EndpointUrl(String),
37    ProviderName(String),
38    Default,
39}
40
41pub struct AdapterRegistry;
42
43impl AdapterRegistry {
44    pub fn select(provider: &ProviderSpec) -> AdapterSelection {
45        if let Some(interface) = provider
46            .interface_hint()
47            .filter(|hint| !hint.starts_with("x_"))
48        {
49            return AdapterSelection::new(
50                provider.api_kind(),
51                AdapterMatchRule::InterfaceHint(interface.to_string()),
52            );
53        }
54
55        let url = provider.endpoints.chat.url.to_ascii_lowercase();
56        if url.contains("/responses") {
57            return AdapterSelection::new(
58                ProviderApi::OpenaiResponses,
59                AdapterMatchRule::EndpointUrl(url),
60            );
61        }
62        if url.contains("/messages") {
63            return AdapterSelection::new(
64                ProviderApi::AnthropicMessagesClaude4,
65                AdapterMatchRule::EndpointUrl(url),
66            );
67        }
68        if provider.provider.eq_ignore_ascii_case("anthropic") {
69            return AdapterSelection::new(
70                ProviderApi::AnthropicMessagesClaude4,
71                AdapterMatchRule::ProviderName(provider.provider.clone()),
72            );
73        }
74
75        AdapterSelection::new(ProviderApi::ChatCompletions, AdapterMatchRule::Default)
76    }
77}
78
79#[cfg(test)]
80mod tests {
81    use super::*;
82    use crate::types::{
83        Capabilities, Constraints, EndpointConfig, Endpoints, HttpMethod, Mappings, ModelConfig,
84        ProviderSpec, RequestMapping, ResponseMapping, SupportFlags,
85    };
86    use std::collections::HashMap;
87
88    fn provider(url: &str, interface: Option<&str>, provider_name: &str) -> ProviderSpec {
89        ProviderSpec {
90            provider: provider_name.into(),
91            models: vec![ModelConfig { id: "m".into() }],
92            interface: interface.map(|s| s.into()),
93            contract_version: Some("1.0.0".into()),
94            inherits: None,
95            endpoints: Endpoints {
96                chat: EndpointConfig {
97                    method: HttpMethod::Post,
98                    url: url.into(),
99                    headers: HashMap::new(),
100                },
101            },
102            mappings: Mappings {
103                request: vec![RequestMapping {
104                    from: "$.messages".into(),
105                    to: "$.messages".into(),
106                    code: None,
107                    clamp: None,
108                }],
109                response: vec![ResponseMapping {
110                    from: "$.choices[0].message".into(),
111                    to: "content".into(),
112                }],
113            },
114            constraints: Constraints {
115                supports: SupportFlags {
116                    json_mode: true,
117                    tools: true,
118                },
119            },
120            auth: crate::auth::AuthScheme::Bearer {
121                token_env: "TOKEN".into(),
122            },
123            capabilities: Capabilities::default(),
124            capabilities_extra: HashMap::new(),
125            extensions: HashMap::new(),
126            unsupported_parameters: Vec::new(),
127        }
128    }
129
130    #[test]
131    fn selects_by_interface_hint() {
132        let spec = provider(
133            "https://api.openai.com/v1/chat/completions",
134            Some("conversational.generate"),
135            "openai",
136        );
137        let selection = AdapterRegistry::select(&spec);
138        assert_eq!(selection.kind(), ProviderApi::ChatCompletions);
139        assert!(matches!(
140            selection.match_rule(),
141            AdapterMatchRule::InterfaceHint(hint) if hint == "conversational.generate"
142        ));
143    }
144
145    #[test]
146    fn selects_by_endpoint_url() {
147        let spec = provider("https://api.openai.com/v1/responses", None, "openai");
148        let selection = AdapterRegistry::select(&spec);
149        assert_eq!(selection.kind(), ProviderApi::OpenaiResponses);
150        assert!(matches!(
151            selection.match_rule(),
152            AdapterMatchRule::EndpointUrl(_)
153        ));
154    }
155
156    #[test]
157    fn selects_by_provider_name() {
158        let spec = provider("https://api.anthropic.com/v1/other", None, "anthropic");
159        let selection = AdapterRegistry::select(&spec);
160        assert_eq!(selection.kind(), ProviderApi::AnthropicMessagesClaude4);
161        assert!(
162            matches!(selection.match_rule(), AdapterMatchRule::ProviderName(name) if name == "anthropic")
163        );
164    }
165
166    #[test]
167    fn falls_back_to_default() {
168        let spec = provider("https://example.com/chat", None, "custom");
169        let selection = AdapterRegistry::select(&spec);
170        assert_eq!(selection.kind(), ProviderApi::ChatCompletions);
171        assert_eq!(selection.match_rule(), &AdapterMatchRule::Default);
172    }
173}