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