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