Skip to main content

gproxy_protocol/protocol/claude/message/
stream.rs

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    /// SSE event name: the wire `type` of this event, if any.
19    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    /// SSE event name: the exact serde rename of this variant.
86    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    /// Regression: `event_name()` must equal the serialized `type` tag.
194    #[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}