Skip to main content

talos_core/
model.rs

1//! Shared model catalog types — provider/model metadata, pricing, and capabilities.
2//!
3//! These types live at the `talos-core` boundary so that multiple crates
4//! (`talos-config`, `talos-models`, CLI, TUI) can share a single definition
5//! without creating cyclic dependencies.
6
7use schemars::JsonSchema;
8use serde::{Deserialize, Serialize};
9
10/// Source of model metadata.
11#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
12#[serde(rename_all = "snake_case")]
13pub enum ModelSource {
14    /// From the built-in dataset embedded at compile time.
15    #[default]
16    Builtin,
17    /// Manually added by the user.
18    Manual,
19    /// Imported from models.dev, with a refresh timestamp.
20    ModelsDev { refreshed_at: String },
21}
22
23/// Pricing information for a model (per 1M tokens, USD).
24#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq)]
25pub struct ModelPricing {
26    /// Input token price per 1M tokens.
27    pub input_per_1m: Option<f64>,
28    /// Output token price per 1M tokens.
29    pub output_per_1m: Option<f64>,
30    /// Cache read token price per 1M tokens.
31    pub cache_read_per_1m: Option<f64>,
32}
33
34/// Capability flags for a model.
35#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
36pub struct ModelCapabilities {
37    /// Supports tool/function calling.
38    #[serde(default)]
39    pub tools: bool,
40    /// Supports structured/JSON output.
41    #[serde(default)]
42    pub structured_output: bool,
43    /// Supports reasoning/thinking mode.
44    #[serde(default)]
45    pub reasoning: bool,
46    /// Accepts image input.
47    #[serde(default)]
48    pub image_input: bool,
49}
50
51/// Image input capability provenance for a model (ADR-050).
52///
53/// `Supported` and `Unsupported` are resolved from confirmed catalog
54/// metadata. `Unknown` applies to custom/discovered models with no
55/// confirmed capability. Both `Unknown` and `Unsupported` fail-closed
56/// for the attachment UI; the distinction is diagnostic.
57///
58/// The default is `Unknown` so that any code path that fails to resolve
59/// a model's metadata cannot accidentally enable image attachment.
60#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
61#[serde(rename_all = "snake_case")]
62pub enum ImageInputCapability {
63    /// Catalog metadata confirms `image_input = true`.
64    Supported,
65    /// Catalog metadata confirms `image_input = false`.
66    Unsupported,
67    /// No confirmed capability (custom/discovered model).
68    #[default]
69    Unknown,
70}
71
72impl ImageInputCapability {
73    /// Resolves the capability from a model's catalog metadata.
74    ///
75    /// Returns `Supported` when `image_input = true`, `Unsupported` when
76    /// `image_input = false`, and `Unknown` when no metadata is available
77    /// (custom/discovered models).
78    pub fn from_metadata(metadata: Option<&ModelMetadata>) -> Self {
79        match metadata {
80            Some(m) if m.capabilities.image_input => Self::Supported,
81            Some(_) => Self::Unsupported,
82            None => Self::Unknown,
83        }
84    }
85
86    /// Returns `true` when image attachment is allowed.
87    pub fn allows_attachment(self) -> bool {
88        matches!(self, Self::Supported)
89    }
90}
91
92/// Reasoning effort levels for OpenAI o-series models.
93#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, JsonSchema)]
94#[serde(rename_all = "lowercase")]
95pub enum ReasoningEffort {
96    /// Low reasoning effort.
97    Low,
98    /// Medium reasoning effort.
99    Medium,
100    /// High reasoning effort.
101    High,
102}
103
104/// Provider API protocol advertised by catalog metadata.
105#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
106pub enum CatalogProviderProtocol {
107    /// Anthropic Messages-compatible API.
108    #[serde(rename = "anthropic-messages")]
109    AnthropicMessages,
110    /// OpenAI Chat Completions-compatible API.
111    #[serde(rename = "openai-chat")]
112    OpenAIChat,
113}
114
115/// Static metadata for a known model.
116///
117/// Represents model knowledge (context limits, pricing, capabilities)
118/// independent of runtime configuration. Used to inform the agent about
119/// model properties without hardcoded fallbacks.
120#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
121pub struct ModelMetadata {
122    /// Unique model identifier (e.g., "claude-sonnet-4-5-20250929").
123    pub id: String,
124    /// Provider name (e.g., "anthropic", "openai", "google").
125    pub provider: String,
126    /// Maximum context window in tokens.
127    pub context_limit: Option<u32>,
128    /// Maximum output tokens.
129    pub output_limit: Option<u32>,
130    /// Pricing information (per 1M tokens, USD).
131    #[serde(default)]
132    pub pricing: Option<ModelPricing>,
133    /// Model capability flags.
134    #[serde(default)]
135    pub capabilities: ModelCapabilities,
136    /// Model release date (ISO 8601 or similar).
137    pub release_date: Option<String>,
138    /// Where this metadata originated.
139    #[serde(default)]
140    pub source: ModelSource,
141    /// Named invocation presets for this model (ADR-048).
142    #[serde(default)]
143    pub variants: Vec<VariantDef>,
144}
145
146/// Look up a model by id in a dataset.
147///
148/// Returns the first model whose `id` matches. When multiple providers share
149/// the same model id (e.g. `glm-5.2` under `zhipu`, `zai`, etc.), this returns
150/// an arbitrary first match. Use [`find_model_by_provider`] for unambiguous
151/// resolution in those cases.
152pub fn find_model<'a>(models: &'a [ModelMetadata], id: &str) -> Option<&'a ModelMetadata> {
153    models.iter().find(|m| m.id == id)
154}
155
156/// Look up a model by `(provider, id)` in a dataset.
157///
158/// Use this instead of [`find_model`] whenever the active provider is known,
159/// so that duplicate model ids across providers resolve unambiguously.
160pub fn find_model_by_provider<'a>(
161    models: &'a [ModelMetadata],
162    provider: &str,
163    id: &str,
164) -> Option<&'a ModelMetadata> {
165    models.iter().find(|m| m.provider == provider && m.id == id)
166}
167
168/// Collects all models whose `id` matches, regardless of provider.
169///
170/// Returns an empty vector when the id is unique or absent. Use this to detect
171/// ambiguity before resolving a bare model id.
172pub fn models_with_id<'a>(models: &'a [ModelMetadata], id: &str) -> Vec<&'a ModelMetadata> {
173    models.iter().filter(|m| m.id == id).collect()
174}
175
176/// Lightweight provider metadata for catalog queries.
177///
178/// Unlike [`talos_config::ProviderConfig`] (which includes credentials and
179/// protocol details), this struct carries only the catalog-level identity and
180/// display information needed by `/model` and `/connect` pickers.
181#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
182pub struct ProviderInfo {
183    /// Provider identifier (e.g., "anthropic", "openai").
184    pub id: String,
185    /// Human-readable display name.
186    pub name: String,
187    /// Default API base URL, if known.
188    pub api_base_url: Option<String>,
189    /// API protocol, if known from catalog metadata.
190    #[serde(default)]
191    pub protocol: Option<CatalogProviderProtocol>,
192    /// Environment variable name for the API key, if conventional.
193    pub env_var: Option<String>,
194    /// Documentation URL, if known.
195    pub doc_url: Option<String>,
196    /// Source of this provider entry.
197    #[serde(default)]
198    pub source: ProviderSource,
199}
200
201/// Source of provider metadata.
202#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
203#[serde(rename_all = "snake_case")]
204pub enum ProviderSource {
205    /// From the built-in dataset.
206    #[default]
207    Builtin,
208    /// Imported from models.dev, with a refresh timestamp.
209    ModelsDev { refreshed_at: String },
210}
211
212#[cfg(test)]
213#[allow(clippy::unwrap_used)]
214mod tests {
215    use super::*;
216
217    #[test]
218    fn test_model_source_default_is_builtin() {
219        assert_eq!(ModelSource::default(), ModelSource::Builtin);
220    }
221
222    #[test]
223    fn test_model_capabilities_default_all_false() {
224        let caps = ModelCapabilities::default();
225        assert!(!caps.tools);
226        assert!(!caps.structured_output);
227        assert!(!caps.reasoning);
228        assert!(!caps.image_input);
229    }
230
231    #[test]
232    fn test_model_pricing_default_all_none() {
233        let pricing = ModelPricing::default();
234        assert!(pricing.input_per_1m.is_none());
235        assert!(pricing.output_per_1m.is_none());
236        assert!(pricing.cache_read_per_1m.is_none());
237    }
238
239    #[test]
240    fn test_find_model_by_provider_resolves() {
241        let models = vec![
242            ModelMetadata {
243                id: "glm-5.2".to_string(),
244                provider: "zhipu".to_string(),
245                context_limit: Some(128_000),
246                output_limit: None,
247                pricing: None,
248                capabilities: ModelCapabilities::default(),
249                release_date: None,
250                variants: vec![],
251                source: ModelSource::Builtin,
252            },
253            ModelMetadata {
254                id: "glm-5.2".to_string(),
255                provider: "zai".to_string(),
256                context_limit: Some(128_000),
257                output_limit: None,
258                pricing: None,
259                capabilities: ModelCapabilities::default(),
260                release_date: None,
261                variants: vec![],
262                source: ModelSource::Builtin,
263            },
264        ];
265
266        let zhipu = find_model_by_provider(&models, "zhipu", "glm-5.2");
267        assert!(zhipu.is_some());
268        assert_eq!(zhipu.expect("operation should succeed").provider, "zhipu");
269
270        let zai = find_model_by_provider(&models, "zai", "glm-5.2");
271        assert!(zai.is_some());
272        assert_eq!(zai.expect("operation should succeed").provider, "zai");
273
274        assert!(find_model_by_provider(&models, "openai", "glm-5.2").is_none());
275    }
276
277    #[test]
278    fn test_models_with_id_detects_ambiguity() {
279        let models = vec![
280            ModelMetadata {
281                id: "shared".to_string(),
282                provider: "a".to_string(),
283                context_limit: None,
284                output_limit: None,
285                pricing: None,
286                capabilities: ModelCapabilities::default(),
287                release_date: None,
288                variants: vec![],
289                source: ModelSource::Builtin,
290            },
291            ModelMetadata {
292                id: "shared".to_string(),
293                provider: "b".to_string(),
294                context_limit: None,
295                output_limit: None,
296                pricing: None,
297                capabilities: ModelCapabilities::default(),
298                release_date: None,
299                variants: vec![],
300                source: ModelSource::Builtin,
301            },
302            ModelMetadata {
303                id: "unique".to_string(),
304                provider: "a".to_string(),
305                context_limit: None,
306                output_limit: None,
307                pricing: None,
308                capabilities: ModelCapabilities::default(),
309                release_date: None,
310                variants: vec![],
311                source: ModelSource::Builtin,
312            },
313        ];
314
315        assert_eq!(models_with_id(&models, "shared").len(), 2);
316        assert_eq!(models_with_id(&models, "unique").len(), 1);
317        assert!(models_with_id(&models, "missing").is_empty());
318    }
319
320    #[test]
321    fn test_provider_info_default() {
322        let info = ProviderInfo::default();
323        assert!(info.id.is_empty());
324        assert!(info.name.is_empty());
325        assert!(info.api_base_url.is_none());
326        assert!(info.protocol.is_none());
327        assert!(info.env_var.is_none());
328        assert!(info.doc_url.is_none());
329        assert_eq!(info.source, ProviderSource::Builtin);
330    }
331
332    #[test]
333    fn test_model_metadata_serde_roundtrip() {
334        let meta = ModelMetadata {
335            id: "test-model".to_string(),
336            provider: "test".to_string(),
337            context_limit: Some(200_000),
338            output_limit: Some(8_192),
339            pricing: Some(ModelPricing {
340                input_per_1m: Some(3.0),
341                output_per_1m: Some(15.0),
342                cache_read_per_1m: Some(0.3),
343            }),
344            capabilities: ModelCapabilities {
345                tools: true,
346                structured_output: false,
347                reasoning: true,
348                image_input: true,
349            },
350            release_date: Some("2025-01-01".to_string()),
351            variants: vec![],
352            source: ModelSource::ModelsDev {
353                refreshed_at: "2025-07-03T00:00:00Z".to_string(),
354            },
355        };
356
357        let json = serde_json::to_string(&meta).expect("serialize");
358        let roundtrip: ModelMetadata = serde_json::from_str(&json).expect("deserialize");
359        assert_eq!(meta.id, roundtrip.id);
360        assert_eq!(meta.provider, roundtrip.provider);
361        assert_eq!(meta.context_limit, roundtrip.context_limit);
362        assert_eq!(meta.output_limit, roundtrip.output_limit);
363        assert_eq!(meta.capabilities, roundtrip.capabilities);
364        assert_eq!(meta.source, roundtrip.source);
365    }
366
367    #[test]
368    fn image_input_capability_supported_when_metadata_image_input_true() {
369        let metadata = ModelMetadata {
370            id: "test-model".into(),
371            provider: "test".into(),
372            context_limit: None,
373            output_limit: None,
374            pricing: None,
375            capabilities: ModelCapabilities {
376                image_input: true,
377                ..Default::default()
378            },
379            release_date: None,
380            source: ModelSource::default(),
381            variants: vec![],
382        };
383        let cap = ImageInputCapability::from_metadata(Some(&metadata));
384        assert_eq!(cap, ImageInputCapability::Supported);
385        assert!(cap.allows_attachment());
386    }
387
388    #[test]
389    fn image_input_capability_unsupported_when_metadata_image_input_false() {
390        let metadata = ModelMetadata {
391            id: "test-model".into(),
392            provider: "test".into(),
393            context_limit: None,
394            output_limit: None,
395            pricing: None,
396            capabilities: ModelCapabilities {
397                image_input: false,
398                ..Default::default()
399            },
400            release_date: None,
401            source: ModelSource::default(),
402            variants: vec![],
403        };
404        let cap = ImageInputCapability::from_metadata(Some(&metadata));
405        assert_eq!(cap, ImageInputCapability::Unsupported);
406        assert!(!cap.allows_attachment());
407    }
408
409    #[test]
410    fn image_input_capability_unknown_when_no_metadata() {
411        let cap = ImageInputCapability::from_metadata(None);
412        assert_eq!(cap, ImageInputCapability::Unknown);
413        assert!(!cap.allows_attachment());
414    }
415}
416
417/// A named invocation preset (ADR-048). Lives in talos-core to avoid
418/// a talos-config dependency from talos-conversation.
419#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, JsonSchema)]
420pub struct VariantDef {
421    /// Stable identifier, e.g. "default", "high-reasoning".
422    pub id: String,
423    /// Display label, e.g. "High Reasoning".
424    pub label: String,
425    /// Optional reasoning effort override.
426    #[serde(default)]
427    pub reasoning_effort: Option<ReasoningEffort>,
428}