Skip to main content

async_llm/
types.rs

1use std::{
2    ops::{Deref, DerefMut},
3    pin::Pin,
4};
5
6use derive_builder::Builder;
7use serde::{Deserialize, Serialize, Serializer};
8use serde_json::Value;
9use tokio_stream::Stream;
10
11use crate::{errors::AnthropicError, messages};
12
13#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, Default)]
14pub struct Usage {
15    pub input_tokens: Option<u32>,
16    pub output_tokens: Option<u32>,
17    /// Tokens written to the prompt cache on this request.
18    /// Populated when one or more `cache_control` markers caused a cache miss.
19    #[serde(default, skip_serializing_if = "Option::is_none")]
20    pub cache_creation_input_tokens: Option<u32>,
21    /// Tokens served from the prompt cache on this request.
22    #[serde(default, skip_serializing_if = "Option::is_none")]
23    pub cache_read_input_tokens: Option<u32>,
24}
25
26/// Marker placed on a content block / tool / system block to define a cache breakpoint.
27///
28/// A breakpoint caches the entire prefix that precedes it (in serialized order:
29/// `tools → system → messages`). Up to four breakpoints per request.
30#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
31#[serde(tag = "type", rename_all = "snake_case")]
32pub enum CacheControl {
33    Ephemeral {
34        /// Cache TTL — `"5m"` (default) or `"1h"`. `None` means the API default.
35        #[serde(default, skip_serializing_if = "Option::is_none")]
36        ttl: Option<String>,
37    },
38}
39
40impl CacheControl {
41    /// Default 5-minute ephemeral breakpoint.
42    #[must_use]
43    pub fn ephemeral() -> Self {
44        CacheControl::Ephemeral { ttl: None }
45    }
46
47    /// Ephemeral breakpoint with an explicit TTL (e.g. `"5m"`, `"1h"`).
48    #[must_use]
49    pub fn ephemeral_with_ttl(ttl: impl Into<String>) -> Self {
50        CacheControl::Ephemeral {
51            ttl: Some(ttl.into()),
52        }
53    }
54}
55
56/// A thinking block returned by extended-thinking models.
57#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
58pub struct Thinking {
59    pub thinking: String,
60    #[serde(default)]
61    pub signature: String,
62    #[serde(default, skip_serializing_if = "Option::is_none")]
63    pub cache_control: Option<CacheControl>,
64}
65
66impl From<Thinking> for MessageContent {
67    fn from(thinking: Thinking) -> Self {
68        MessageContent::Thinking(thinking)
69    }
70}
71
72impl From<Thinking> for MessageContentList {
73    fn from(thinking: Thinking) -> Self {
74        MessageContentList(vec![thinking.into()])
75    }
76}
77
78/// Configuration for extended thinking.
79#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
80#[serde(tag = "type", rename_all = "snake_case")]
81pub enum ThinkingConfig {
82    Enabled { budget_tokens: u32 },
83    Disabled,
84}
85
86#[derive(Clone, Debug, Deserialize)]
87pub enum ToolChoice {
88    Auto,
89    Any,
90    Tool(String),
91}
92
93#[derive(Debug, Clone, Serialize, Deserialize, Builder, PartialEq, Default)]
94#[builder(setter(into, strip_option), default)]
95pub struct Message {
96    pub role: MessageRole,
97    pub content: MessageContentList,
98}
99
100impl Message {
101    /// Returns all the tool uses in the message
102    pub fn tool_uses(&self) -> Vec<ToolUse> {
103        self.content
104            .0
105            .iter()
106            .filter(|c| matches!(c, MessageContent::ToolUse(_)))
107            .map(|c| match c {
108                MessageContent::ToolUse(tool_use) => tool_use.clone(),
109                _ => unreachable!(),
110            })
111            .collect()
112    }
113
114    /// Returns the first text content in the message
115    pub fn text(&self) -> Option<String> {
116        self.content
117            .0
118            .iter()
119            .filter_map(|c| match c {
120                MessageContent::Text(text) => Some(text.text.clone()),
121                _ => None,
122            })
123            .next()
124    }
125}
126
127#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
128pub struct MessageContentList(pub Vec<MessageContent>);
129
130impl Deref for MessageContentList {
131    type Target = Vec<MessageContent>;
132
133    fn deref(&self) -> &Self::Target {
134        &self.0
135    }
136}
137
138impl DerefMut for MessageContentList {
139    fn deref_mut(&mut self) -> &mut Self::Target {
140        &mut self.0
141    }
142}
143
144#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
145#[serde(rename_all = "snake_case")]
146pub enum MessageRole {
147    #[default]
148    User,
149    Assistant,
150}
151
152#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
153#[builder(setter(into, strip_option))]
154pub struct CreateMessagesRequest {
155    pub messages: Vec<Message>,
156    pub model: String,
157    #[builder(default = messages::DEFAULT_MAX_TOKENS)]
158    pub max_tokens: i32,
159    #[builder(default)]
160    #[serde(skip_serializing_if = "Option::is_none")]
161    pub metadata: Option<serde_json::Map<String, Value>>,
162    #[serde(skip_serializing_if = "Option::is_none")]
163    #[builder(default)]
164    pub stop_sequences: Option<Vec<String>>,
165    #[builder(default = "false")]
166    pub stream: bool, // Optional default false
167    #[serde(skip_serializing_if = "Option::is_none")]
168    #[builder(default)]
169    pub temperature: Option<f32>, // 0 < x < 1
170    #[serde(skip_serializing_if = "Option::is_none")]
171    #[builder(default)]
172    pub tool_choice: Option<ToolChoice>,
173    // TODO: Type this
174    #[serde(skip_serializing_if = "Option::is_none")]
175    #[builder(default)]
176    pub tools: Option<Vec<serde_json::Map<String, Value>>>,
177    #[serde(skip_serializing_if = "Option::is_none")]
178    #[builder(default)]
179    pub top_k: Option<u32>, // > 0
180    #[serde(skip_serializing_if = "Option::is_none")]
181    #[builder(default)]
182    pub top_p: Option<f32>, // 0 < x < 1
183    #[serde(skip_serializing_if = "Option::is_none")]
184    #[builder(default)]
185    pub system: Option<String>,
186    #[serde(skip_serializing_if = "Option::is_none")]
187    #[builder(default)]
188    pub thinking: Option<ThinkingConfig>,
189}
190
191#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
192#[builder(setter(into, strip_option))]
193pub struct CreateMessagesResponse {
194    #[serde(default)]
195    pub id: Option<String>,
196    #[serde(default)]
197    pub content: Option<Vec<MessageContent>>,
198    #[serde(default)]
199    pub model: Option<String>,
200    #[serde(default)]
201    pub stop_reason: Option<String>,
202    #[serde(default)]
203    pub stop_sequence: Option<String>,
204    #[serde(default)]
205    pub usage: Option<Usage>,
206}
207
208impl CreateMessagesResponse {
209    /// Returns the content as Messages so they are more easily reusable
210    pub fn messages(&self) -> Vec<Message> {
211        let Some(content) = &self.content else {
212            return vec![];
213        };
214        content
215            .iter()
216            .map(|c| Message {
217                role: MessageRole::Assistant,
218                content: c.clone().into(),
219            })
220            .collect()
221    }
222}
223
224#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
225#[serde(tag = "type", rename_all = "snake_case")]
226pub enum MessageContent {
227    ToolUse(ToolUse),
228    ToolResult(ToolResult),
229    Text(Text),
230    Thinking(Thinking),
231    // TODO: Implement images and documents
232}
233
234impl MessageContent {
235    pub fn as_tool_use(&self) -> Option<&ToolUse> {
236        if let MessageContent::ToolUse(tool_use) = self {
237            Some(tool_use)
238        } else {
239            None
240        }
241    }
242
243    pub fn as_tool_result(&self) -> Option<&ToolResult> {
244        if let MessageContent::ToolResult(tool_result) = self {
245            Some(tool_result)
246        } else {
247            None
248        }
249    }
250
251    pub fn as_text(&self) -> Option<&Text> {
252        if let MessageContent::Text(text) = self {
253            Some(text)
254        } else {
255            None
256        }
257    }
258
259    pub fn as_thinking(&self) -> Option<&Thinking> {
260        if let MessageContent::Thinking(thinking) = self {
261            Some(thinking)
262        } else {
263            None
264        }
265    }
266}
267
268#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default, Builder)]
269#[builder(setter(into, strip_option), default)]
270pub struct ToolUse {
271    pub id: String,
272    pub input: Value,
273    pub name: String,
274    #[serde(default, skip_serializing_if = "Option::is_none")]
275    pub cache_control: Option<CacheControl>,
276}
277
278impl From<ToolUse> for MessageContent {
279    fn from(tool_use: ToolUse) -> Self {
280        MessageContent::ToolUse(tool_use)
281    }
282}
283
284impl From<ToolUse> for MessageContentList {
285    fn from(tool_use: ToolUse) -> Self {
286        MessageContentList(vec![tool_use.into()])
287    }
288}
289
290#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default, Builder)]
291#[builder(setter(into, strip_option), default)]
292pub struct ToolResult {
293    pub tool_use_id: String,
294    pub content: Option<String>,
295    pub is_error: bool,
296    #[serde(default, skip_serializing_if = "Option::is_none")]
297    pub cache_control: Option<CacheControl>,
298}
299
300impl From<ToolResult> for MessageContent {
301    fn from(tool_result: ToolResult) -> Self {
302        MessageContent::ToolResult(tool_result)
303    }
304}
305
306impl From<ToolResult> for MessageContentList {
307    fn from(tool_result: ToolResult) -> Self {
308        MessageContentList(vec![tool_result.into()])
309    }
310}
311
312#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default, Builder)]
313#[builder(setter(into, strip_option), default)]
314pub struct Text {
315    pub text: String,
316    #[serde(default, skip_serializing_if = "Option::is_none")]
317    pub cache_control: Option<CacheControl>,
318}
319
320impl<S: AsRef<str>> From<S> for Text {
321    fn from(s: S) -> Self {
322        Text {
323            text: s.as_ref().to_string(),
324            ..Default::default()
325        }
326    }
327}
328
329impl From<Text> for MessageContent {
330    fn from(text: Text) -> Self {
331        MessageContent::Text(text)
332    }
333}
334
335impl From<Text> for MessageContentList {
336    fn from(text: Text) -> Self {
337        MessageContentList(vec![text.into()])
338    }
339}
340
341impl<S: AsRef<str>> From<S> for MessageContent {
342    fn from(s: S) -> Self {
343        MessageContent::Text(Text {
344            text: s.as_ref().to_string(),
345            ..Default::default()
346        })
347    }
348}
349
350impl<S: AsRef<str>> From<S> for Message {
351    fn from(s: S) -> Self {
352        MessageBuilder::default()
353            .role(MessageRole::User)
354            .content(s.as_ref().to_string())
355            .build()
356            .expect("infallible")
357    }
358}
359
360// Any single AsRef<str> can be converted to a MessageContent, in a list as a single item
361impl<S: AsRef<str>> From<S> for MessageContentList {
362    fn from(s: S) -> Self {
363        MessageContentList(vec![s.as_ref().into()])
364    }
365}
366
367impl From<MessageContent> for MessageContentList {
368    fn from(content: MessageContent) -> Self {
369        MessageContentList(vec![content])
370    }
371}
372
373impl Serialize for ToolChoice {
374    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
375    where
376        S: Serializer,
377    {
378        match self {
379            ToolChoice::Auto => {
380                serde::Serialize::serialize(&serde_json::json!({"type": "auto"}), serializer)
381            }
382            ToolChoice::Any => {
383                serde::Serialize::serialize(&serde_json::json!({"type": "any"}), serializer)
384            }
385            ToolChoice::Tool(name) => serde::Serialize::serialize(
386                &serde_json::json!({"type": "tool", "name": name}),
387                serializer,
388            ),
389        }
390    }
391}
392#[derive(Clone, Serialize, Deserialize, Debug, Eq, PartialEq)]
393#[serde(rename_all = "snake_case", tag = "type")]
394pub enum ContentBlockDelta {
395    TextDelta { text: String },
396    InputJsonDelta { partial_json: String },
397    ThinkingDelta { thinking: String },
398    SignatureDelta { signature: String },
399}
400
401#[derive(Clone, Serialize, Deserialize, Debug, Eq, PartialEq)]
402pub struct MessageDelta {
403    pub stop_reason: Option<String>,
404    pub stop_sequence: Option<String>,
405}
406
407#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
408#[serde(rename_all = "snake_case", tag = "type")]
409pub enum MessagesStreamEvent {
410    MessageStart {
411        message: MessageStart,
412        usage: Option<Usage>,
413    },
414    ContentBlockStart {
415        index: usize,
416        content_block: MessageContent,
417    },
418    ContentBlockDelta {
419        index: usize,
420        delta: ContentBlockDelta,
421    },
422    ContentBlockStop {
423        index: usize,
424    },
425    MessageDelta {
426        delta: MessageDelta,
427        #[serde(default)]
428        usage: Option<Usage>,
429    },
430    MessageStop,
431}
432#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
433pub struct MessageStart {
434    pub id: String,
435    pub model: String,
436    pub role: String,
437    pub content: Vec<MessageContent>,
438    #[serde(default)]
439    pub stop_reason: Option<String>,
440    #[serde(default)]
441    pub stop_sequence: Option<String>,
442    #[serde(default)]
443    pub usage: Option<Usage>,
444}
445
446pub type CreateMessagesResponseStream =
447    Pin<Box<dyn Stream<Item = Result<MessagesStreamEvent, AnthropicError>> + Send>>;
448
449#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
450pub struct ListModelsResponse {
451    #[serde(default)]
452    pub data: Vec<Model>,
453
454    #[serde(default)]
455    pub first_id: Option<String>,
456    pub has_more: bool,
457    #[serde(default)]
458    pub last_id: Option<String>,
459}
460
461#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
462pub struct Model {
463    pub created_at: String,
464    pub display_name: String,
465    pub id: String,
466    #[serde(rename = "type")]
467    pub model_type: String,
468}
469
470pub type GetModelResponse = Model;
471
472#[cfg(test)]
473mod tests {
474    use serde_json::json;
475
476    use super::*;
477
478    #[test_log::test(tokio::test)]
479    async fn test_deserialize_response() {
480        let response = json!({
481        "id":"msg_01KkaCASJuaAgTWD2wqdbwC8",
482        "type":"message",
483        "role":"assistant",
484        "model":"claude-3-5-sonnet-20241022",
485        "content":[
486            {"type":"text",
487        "text":"Hi! How can I help you today?"}],
488        "stop_reason":"end_turn",
489        "stop_sequence":null,
490        "usage":{"input_tokens":10,"cache_creation_input_tokens":0,"cache_read_input_tokens":0,"output_tokens":12}}).to_string();
491
492        let response = serde_json::from_str::<CreateMessagesResponse>(&response).unwrap();
493
494        let usage = response.usage.as_ref().unwrap();
495
496        assert_eq!(usage.input_tokens, Some(10));
497        assert_eq!(usage.output_tokens, Some(12));
498        assert_eq!(usage.cache_creation_input_tokens, Some(0));
499        assert_eq!(usage.cache_read_input_tokens, Some(0));
500        assert_eq!(
501            response.id,
502            Some("msg_01KkaCASJuaAgTWD2wqdbwC8".to_string())
503        );
504        assert_eq!(
505            response.model,
506            Some("claude-3-5-sonnet-20241022".to_string())
507        );
508        assert_eq!(response.stop_reason, Some("end_turn".to_string()));
509        assert_eq!(response.stop_sequence, None);
510        assert_eq!(
511            response
512                .messages()
513                .first()
514                .unwrap()
515                .content
516                .first()
517                .unwrap()
518                .as_text(),
519            Some(&Text {
520                text: "Hi! How can I help you today?".to_string(),
521                cache_control: None,
522            })
523        );
524    }
525
526    #[test_log::test(tokio::test)]
527    async fn test_from_str() {
528        let message: Message = "Hello world!".into();
529
530        assert_eq!(
531            message,
532            Message {
533                role: MessageRole::User,
534                content: MessageContentList(vec![MessageContent::Text(Text {
535                    text: "Hello world!".to_string(),
536                    cache_control: None,
537                })]),
538            }
539        );
540
541        assert_eq!(message.text(), Some("Hello world!".to_string()));
542    }
543}