Skip to main content

rig_core/providers/anthropic/
modality.rs

1//! Model listing and credential verification for a Messages-format
2//! provider: `GET /v1/models`, paged by cursor, and the same endpoint read
3//! for its status alone.
4//!
5//! ```no_run
6//! use rig_core::providers::anthropic::Anthropic;
7//!
8//! # async fn run() -> Result<(), Box<dyn std::error::Error>> {
9//! let models = Anthropic::from_env()?.list_models().await?;
10//! # let _ = models;
11//! # Ok(())
12//! # }
13//! ```
14
15use serde::{Deserialize, Serialize};
16
17use super::wire::AnthropicConfig;
18use crate::error::{EncodeError, ProviderError};
19use crate::model::{ModelInfo, ModelList};
20pub use crate::operation::VerifyDecoder;
21use crate::operation::{ModelListing, ModelPage, Verify as VerifyOp};
22use crate::wire::{
23    Body, Decoder, Descriptor, Encoded, Flow, Framing, Mode, Out, Wire, WireEvent, WireFrame,
24};
25
26impl AnthropicConfig {
27    /// The model-listing wire.
28    pub(crate) fn models(&self) -> Models {
29        Models {
30            provider: self.clone(),
31        }
32    }
33
34    /// The credential-check wire.
35    pub(crate) fn verify(&self) -> Verify {
36        Verify {
37            provider: self.clone(),
38        }
39    }
40}
41
42/// The model-listing wire: `GET /v1/models`, cursor-paged.
43#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
44pub struct Models {
45    /// The provider this wire speaks to.
46    pub provider: AnthropicConfig,
47}
48
49impl Wire for Models {
50    type Op = ModelListing;
51    type Payload = crate::wire::Encoded;
52    type Frame = crate::wire::WireFrame;
53    type Decoder<'id> = ModelsDecoder;
54    type Reassembler = crate::wire::document::Unreassembled;
55
56    fn describe(&self) -> Descriptor<'_> {
57        Descriptor::new(self.provider.dialect.name)
58    }
59
60    fn encode(&self, cursor: Option<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
61        Ok(Encoded::new(
62            self.models_request(cursor.as_deref())?,
63            Framing::Whole,
64        ))
65    }
66
67    fn decoder<'id>(&self) -> Self::Decoder<'id> {
68        ModelsDecoder
69    }
70}
71
72impl Models {
73    /// One page's request, after `cursor` when the previous page named one.
74    fn models_request(&self, cursor: Option<&str>) -> Result<http::Request<Body>, EncodeError> {
75        let uri = match cursor {
76            Some(cursor) => format!(
77                "{}{}",
78                self.provider.base_url,
79                crate::providers::internal::with_query_pairs("/v1/models", &[("after_id", cursor)],)
80            ),
81            None => format!("{}/v1/models", self.provider.base_url),
82        };
83        self.provider
84            .headers(http::Request::get(uri))
85            .body(Body::empty())
86            .map_err(EncodeError::from)
87    }
88}
89
90/// One page of `GET /v1/models`.
91#[derive(Debug, Deserialize)]
92#[doc(hidden)]
93pub struct ModelsPage {
94    data: Vec<ModelEntry>,
95    #[serde(default)]
96    has_more: bool,
97    #[serde(default)]
98    last_id: Option<String>,
99}
100
101#[derive(Debug, Deserialize)]
102struct ModelEntry {
103    id: String,
104    display_name: String,
105}
106
107impl From<ModelEntry> for ModelInfo {
108    fn from(entry: ModelEntry) -> Self {
109        ModelInfo::new(entry.id, entry.display_name)
110    }
111}
112
113/// Decodes `GET /v1/models` and the cursor Anthropic names.
114pub struct ModelsDecoder;
115
116impl<'id> Decoder<'id, ModelListing> for ModelsDecoder {
117    type Event = ModelsPage;
118
119    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
120        crate::providers::internal::wire::classify_marker_keyed_frame(&frame.as_str(), &["data"])
121    }
122
123    fn decode(
124        &mut self,
125        page: Self::Event,
126        out: Out<'id, ModelListing>,
127    ) -> Result<Flow, ProviderError> {
128        // Missing or empty cursors would repeatedly fetch page one, even with has_more.
129        let next = page
130            .last_id
131            .filter(|cursor| page.has_more && !cursor.is_empty());
132        Ok(out.end(ModelPage {
133            models: ModelList::new(page.data.into_iter().map(ModelInfo::from).collect()),
134            next,
135        }))
136    }
137}
138
139/// The credential-check wire: `GET /v1/models`, status only, decoded by
140/// the shared [`VerifyDecoder`].
141#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
142pub struct Verify {
143    /// The provider this wire speaks to.
144    pub provider: AnthropicConfig,
145}
146
147impl Wire for Verify {
148    type Op = VerifyOp;
149    type Payload = crate::wire::Encoded;
150    type Frame = crate::wire::WireFrame;
151    type Decoder<'id> = VerifyDecoder;
152    type Reassembler = crate::wire::document::Unreassembled;
153
154    fn describe(&self) -> Descriptor<'_> {
155        Descriptor::new(self.provider.dialect.name)
156    }
157
158    fn encode(&self, _request: (), _mode: Mode) -> Result<Encoded, EncodeError> {
159        let request = self
160            .provider
161            .headers(http::Request::get(format!(
162                "{}/v1/models",
163                self.provider.base_url
164            )))
165            .body(Body::empty())?;
166        Ok(Encoded::new(request, Framing::Whole))
167    }
168
169    fn decoder<'id>(&self) -> Self::Decoder<'id> {
170        VerifyDecoder
171    }
172}