1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4
5use crate::protocol::claude::common::{
6 AssistantRole, Citation, ClaudeModel, Container, ContentBlock, ContextManagementResponse,
7 JsonObject, MessageObjectType, StopDetails, StopReason, TypedObject, Usage,
8};
9
10#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
11#[serde(untagged)]
12pub enum StreamEvent {
13 Known(Box<KnownStreamEvent>),
14 Unknown(TypedObject),
15}
16
17impl StreamEvent {
18 pub fn event_name(&self) -> Option<&str> {
20 match self {
21 Self::Known(event) => Some(event.event_name()),
22 Self::Unknown(object) => Some(object.type_.as_str()),
23 }
24 }
25}
26
27#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28#[serde(tag = "type")]
29pub enum KnownStreamEvent {
30 #[serde(rename = "message_start")]
31 MessageStart {
32 message: Box<CreateMessageStartBody>,
33 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
34 extra: JsonObject,
35 },
36 #[serde(rename = "content_block_start")]
37 ContentBlockStart {
38 index: u64,
39 content_block: Box<ContentBlock>,
40 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
41 extra: JsonObject,
42 },
43 #[serde(rename = "content_block_delta")]
44 ContentBlockDelta {
45 index: u64,
46 delta: Box<EventDelta>,
47 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
48 extra: JsonObject,
49 },
50 #[serde(rename = "content_block_stop")]
51 ContentBlockStop {
52 index: u64,
53 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
54 extra: JsonObject,
55 },
56 #[serde(rename = "message_delta")]
57 MessageDelta {
58 #[serde(skip_serializing_if = "Option::is_none")]
59 context_management: Option<Box<ContextManagementResponse>>,
60 delta: Box<MessageDelta>,
61 #[serde(skip_serializing_if = "Option::is_none")]
62 usage: Option<Box<Usage>>,
63 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
64 extra: JsonObject,
65 },
66 #[serde(rename = "message_stop")]
67 MessageStop {
68 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
69 extra: JsonObject,
70 },
71 #[serde(rename = "ping")]
72 Ping {
73 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
74 extra: JsonObject,
75 },
76 #[serde(rename = "error")]
77 Error {
78 error: StreamError,
79 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
80 extra: JsonObject,
81 },
82}
83
84impl KnownStreamEvent {
85 pub fn event_name(&self) -> &'static str {
87 match self {
88 Self::MessageStart { .. } => "message_start",
89 Self::ContentBlockStart { .. } => "content_block_start",
90 Self::ContentBlockDelta { .. } => "content_block_delta",
91 Self::ContentBlockStop { .. } => "content_block_stop",
92 Self::MessageDelta { .. } => "message_delta",
93 Self::MessageStop { .. } => "message_stop",
94 Self::Ping { .. } => "ping",
95 Self::Error { .. } => "error",
96 }
97 }
98}
99
100#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
101#[serde(untagged)]
102pub enum EventDelta {
103 Known(Box<KnownEventDelta>),
104 Unknown(TypedObject),
105}
106
107#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
108#[serde(tag = "type")]
109pub enum KnownEventDelta {
110 #[serde(rename = "text_delta")]
111 Text {
112 text: String,
113 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
114 extra: JsonObject,
115 },
116 #[serde(rename = "input_json_delta")]
117 InputJson {
118 partial_json: String,
119 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
120 extra: JsonObject,
121 },
122 #[serde(rename = "citations_delta")]
123 Citations {
124 citation: Box<Citation>,
125 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
126 extra: JsonObject,
127 },
128 #[serde(rename = "thinking_delta")]
129 Thinking {
130 #[serde(default, skip_serializing_if = "Option::is_none")]
131 estimated_tokens: Option<u64>,
132 thinking: String,
133 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
134 extra: JsonObject,
135 },
136 #[serde(rename = "signature_delta")]
137 Signature {
138 signature: String,
139 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
140 extra: JsonObject,
141 },
142 #[serde(rename = "compaction_delta")]
143 Compaction {
144 content: String,
145 encrypted_content: String,
146 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
147 extra: JsonObject,
148 },
149}
150
151#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
152pub struct CreateMessageStartBody {
153 pub id: String,
154 #[serde(rename = "type")]
155 pub type_: MessageObjectType,
156 pub role: AssistantRole,
157 pub content: Vec<ContentBlock>,
158 pub model: ClaudeModel,
159 pub stop_reason: Option<StopReason>,
160 pub stop_sequence: Option<String>,
161 pub usage: Usage,
162 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
163 pub extra: JsonObject,
164}
165
166#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
167pub struct MessageDelta {
168 #[serde(skip_serializing_if = "Option::is_none")]
169 pub container: Option<Container>,
170 #[serde(skip_serializing_if = "Option::is_none")]
171 pub stop_reason: Option<StopReason>,
172 #[serde(skip_serializing_if = "Option::is_none")]
173 pub stop_sequence: Option<String>,
174 #[serde(skip_serializing_if = "Option::is_none")]
175 pub stop_details: Option<StopDetails>,
176 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
177 pub extra: JsonObject,
178}
179
180#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
181pub struct StreamError {
182 #[serde(rename = "type")]
183 pub type_: String,
184 pub message: String,
185 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
186 pub extra: JsonObject,
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 #[test]
195 fn event_name_matches_serialized_type_tag() {
196 let events: Vec<StreamEvent> = [
197 r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}"#,
198 r#"{"type":"message_stop"}"#,
199 r#"{"type":"some_future_event","x":1}"#,
200 ]
201 .iter()
202 .map(|raw| serde_json::from_str(raw).unwrap())
203 .collect();
204 for event in events {
205 let value = serde_json::to_value(&event).unwrap();
206 assert_eq!(
207 event.event_name(),
208 value.get("type").and_then(serde_json::Value::as_str),
209 "{value}"
210 );
211 }
212 }
213}