Skip to main content

sim_lib_openai_server/codec_openai/
streaming.rs

1use serde_json::{Value, json};
2use sim_kernel::{Error, Expr, Result, Symbol};
3use sim_lib_stream_core::StreamPacket;
4
5use crate::{codec_openai::encode_openai_responses_response, objects::GatewayEvent};
6
7/// Selects which OpenAI server-sent-event surface to render gateway events as.
8#[derive(Clone, Copy, Debug, PartialEq, Eq)]
9pub enum OpenAiSseSurface {
10    /// The OpenAI Responses API streaming surface.
11    Responses,
12    /// The OpenAI Chat Completions streaming surface.
13    Chat,
14}
15
16/// Decoded gateway event payload: sequence number, event kind, and body.
17#[derive(Clone, Debug, PartialEq, Eq)]
18pub struct GatewayEventData {
19    sequence: u64,
20    kind: Symbol,
21    payload: Expr,
22}
23
24impl GatewayEventData {
25    /// Extracts the streaming-relevant fields from a [`GatewayEvent`].
26    pub fn from_event(event: &GatewayEvent) -> Self {
27        Self {
28            sequence: event.sequence(),
29            kind: event.kind().clone(),
30            payload: event.payload().clone(),
31        }
32    }
33
34    /// Decodes gateway event data from a stream data packet, erroring if the
35    /// packet is not a data packet of the gateway event kind.
36    pub fn from_packet(packet: &StreamPacket) -> Result<Self> {
37        let StreamPacket::Data(data) = packet else {
38            return Err(Error::TypeMismatch {
39                expected: "OpenAI gateway data packet",
40                found: "non-data packet",
41            });
42        };
43        if data.kind != gateway_event_data_kind() {
44            return Err(Error::Eval(format!(
45                "expected OpenAI gateway data kind {}, found {}",
46                gateway_event_data_kind(),
47                data.kind
48            )));
49        }
50        Ok(Self {
51            sequence: string_field(&data.payload, "sequence")?
52                .parse::<u64>()
53                .map_err(|err| Error::Eval(format!("invalid gateway event sequence: {err}")))?,
54            kind: symbol_value_field(&data.payload, "event-kind")?,
55            payload: map_field(&data.payload, "payload")?.clone(),
56        })
57    }
58
59    /// Returns the event sequence number.
60    pub fn sequence(&self) -> u64 {
61        self.sequence
62    }
63
64    /// Returns the event kind symbol.
65    pub fn kind(&self) -> &Symbol {
66        &self.kind
67    }
68
69    /// Returns the event payload expression.
70    pub fn payload(&self) -> &Expr {
71        &self.payload
72    }
73}
74
75/// Returns the stream data kind symbol `stream/data:openai-gateway-event`.
76pub fn gateway_event_data_kind() -> Symbol {
77    Symbol::qualified("stream/data", "openai-gateway-event")
78}
79
80/// Wraps gateway events as stream data packets of the gateway event kind.
81pub fn gateway_event_data_packets(events: &[GatewayEvent]) -> Vec<StreamPacket> {
82    events
83        .iter()
84        .map(|event| StreamPacket::data(gateway_event_data_kind(), event.to_expr()))
85        .collect()
86}
87
88/// Decodes [`GatewayEventData`] from a stream data packet.
89pub fn gateway_event_data_from_packet(packet: &StreamPacket) -> Result<GatewayEventData> {
90    GatewayEventData::from_packet(packet)
91}
92
93/// Accumulates server-sent-event (SSE) framed bytes for a streamed response.
94#[derive(Clone, Debug, Default)]
95pub struct StreamSink {
96    body: Vec<u8>,
97    done: bool,
98}
99
100impl StreamSink {
101    /// Creates an empty sink.
102    pub fn new() -> Self {
103        Self::default()
104    }
105
106    /// Appends one `data:` SSE frame carrying the JSON `value`.
107    pub fn event(&mut self, value: Value) -> Result<()> {
108        let line = crate::objects::canonical_json_bytes(value);
109        self.body.extend_from_slice(b"data: ");
110        self.body.extend_from_slice(&line);
111        self.body.extend_from_slice(b"\n\n");
112        Ok(())
113    }
114
115    /// Appends the terminating `data: [DONE]` frame once, if not already done.
116    pub fn done(&mut self) {
117        if !self.done {
118            self.body.extend_from_slice(b"data: [DONE]\n\n");
119            self.done = true;
120        }
121    }
122
123    /// Finalizes the stream (emitting `[DONE]`) and returns the framed bytes.
124    pub fn into_bytes(mut self) -> Vec<u8> {
125        self.done();
126        self.body
127    }
128}
129
130/// Encodes a sequence of gateway events as SSE bytes for the given surface,
131/// using `response_id` and `created_at_ms` to stamp the emitted chunks.
132pub fn encode_gateway_events_sse(
133    events: &[GatewayEvent],
134    surface: OpenAiSseSurface,
135    response_id: &str,
136    created_at_ms: u64,
137) -> Result<Vec<u8>> {
138    let packets = gateway_event_data_packets(events);
139    let events = packets
140        .iter()
141        .map(gateway_event_data_from_packet)
142        .collect::<Result<Vec<_>>>()?;
143    let model = model_from_event_data(&events).unwrap_or_else(|| "fixture/echo".to_owned());
144    let mut sink = StreamSink::new();
145    match surface {
146        OpenAiSseSurface::Responses => {
147            for event in &events {
148                if let Some(chunk) = responses_chunk(event, response_id, created_at_ms, &model)? {
149                    sink.event(chunk)?;
150                }
151            }
152        }
153        OpenAiSseSurface::Chat => {
154            for event in &events {
155                if let Some(chunk) = chat_chunk(event, response_id, created_at_ms, &model)? {
156                    sink.event(chunk)?;
157                }
158            }
159        }
160    }
161    Ok(sink.into_bytes())
162}
163
164fn responses_chunk(
165    event: &GatewayEventData,
166    response_id: &str,
167    created_at_ms: u64,
168    model: &str,
169) -> Result<Option<Value>> {
170    Ok(match event.kind().name.as_ref() {
171        "request-start" => Some(json!({
172            "type": "response.created",
173            "response": response_stub(response_id, created_at_ms, model, "created"),
174        })),
175        "plan-start" => Some(json!({
176            "type": "response.metadata",
177            "sequence": event.sequence(),
178        })),
179        "model-start" => Some(json!({
180            "type": "response.in_progress",
181            "response": response_stub(response_id, created_at_ms, model, "in_progress"),
182        })),
183        "delta" => Some(json!({
184            "type": "response.output_text.delta",
185            "delta": string_payload(event.payload())?,
186        })),
187        "usage" => Some(json!({
188            "type": "response.usage",
189            "usage": usage_json(event.payload())?,
190        })),
191        "error" => Some(json!({
192            "type": "error",
193            "error": error_json(event.payload()),
194        })),
195        "final" => Some(json!({
196            "type": "response.completed",
197            "response": final_response_json(event.payload(), response_id, created_at_ms)?,
198        })),
199        _ => None,
200    })
201}
202
203fn chat_chunk(
204    event: &GatewayEventData,
205    response_id: &str,
206    created_at_ms: u64,
207    model: &str,
208) -> Result<Option<Value>> {
209    Ok(match event.kind().name.as_ref() {
210        "model-start" => Some(chat_choice_chunk(
211            response_id,
212            created_at_ms,
213            model,
214            json!({"role": "assistant"}),
215            Value::Null,
216        )),
217        "delta" => Some(chat_choice_chunk(
218            response_id,
219            created_at_ms,
220            model,
221            json!({"content": string_payload(event.payload())?}),
222            Value::Null,
223        )),
224        "usage" => Some(json!({
225            "id": response_id,
226            "object": "chat.completion.chunk",
227            "created": created_at_ms / 1000,
228            "model": model,
229            "choices": [],
230            "usage": usage_json(event.payload())?,
231        })),
232        "error" => Some(json!({
233            "type": "error",
234            "error": error_json(event.payload()),
235        })),
236        "final" => Some(chat_choice_chunk(
237            response_id,
238            created_at_ms,
239            model,
240            json!({}),
241            json!(finish_reason(event.payload()).unwrap_or_else(|_| "stop".to_owned())),
242        )),
243        _ => None,
244    })
245}
246
247fn response_stub(response_id: &str, created_at_ms: u64, model: &str, status: &str) -> Value {
248    json!({
249        "id": response_id,
250        "object": "response",
251        "created_at": created_at_ms / 1000,
252        "status": status,
253        "model": model,
254    })
255}
256
257fn chat_choice_chunk(
258    response_id: &str,
259    created_at_ms: u64,
260    model: &str,
261    delta: Value,
262    finish_reason: Value,
263) -> Value {
264    json!({
265        "id": response_id,
266        "object": "chat.completion.chunk",
267        "created": created_at_ms / 1000,
268        "model": model,
269        "choices": [{
270            "index": 0,
271            "delta": delta,
272            "finish_reason": finish_reason,
273        }],
274    })
275}
276
277fn final_response_json(expr: &Expr, response_id: &str, created_at_ms: u64) -> Result<Value> {
278    let bytes = encode_openai_responses_response(expr, response_id, created_at_ms)?;
279    serde_json::from_slice(&bytes).map_err(|err| {
280        Error::Eval(format!(
281            "openai codec failed to decode final response chunk: {err}"
282        ))
283    })
284}
285
286fn model_from_event_data(events: &[GatewayEventData]) -> Option<String> {
287    events.iter().find_map(|event| {
288        if event.kind().name.as_ref() == "model-start" {
289            string_payload(event.payload()).ok()
290        } else if event.kind().name.as_ref() == "final" {
291            string_field(event.payload(), "model").ok()
292        } else {
293            None
294        }
295    })
296}
297
298fn finish_reason(expr: &Expr) -> Result<String> {
299    symbol_field(expr, "stop-reason")
300}
301
302fn usage_json(expr: &Expr) -> Result<Value> {
303    let Expr::Map(fields) = expr else {
304        return Err(Error::Eval(
305            "openai SSE usage payload must be a map".to_owned(),
306        ));
307    };
308    let prompt = optional_u64_field(fields, "input-tokens")?.unwrap_or(0);
309    let completion = optional_u64_field(fields, "output-tokens")?.unwrap_or(0);
310    let total = optional_u64_field(fields, "total-tokens")?.unwrap_or(prompt + completion);
311    Ok(json!({
312        "prompt_tokens": prompt,
313        "completion_tokens": completion,
314        "total_tokens": total,
315    }))
316}
317
318fn error_json(expr: &Expr) -> Value {
319    match expr {
320        Expr::String(message) => json!({"message": message}),
321        other => json!({"message": format!("{other:?}")}),
322    }
323}
324
325fn string_payload(expr: &Expr) -> Result<String> {
326    match expr {
327        Expr::String(text) => Ok(text.clone()),
328        other => Err(Error::Eval(format!(
329            "openai SSE event payload must be a string, found {other:?}"
330        ))),
331    }
332}
333
334fn string_field(expr: &Expr, key: &str) -> Result<String> {
335    match map_field(expr, key)? {
336        Expr::String(text) => Ok(text.clone()),
337        _ => Err(Error::Eval(format!(
338            "openai SSE field {key} must be a string"
339        ))),
340    }
341}
342
343fn symbol_field(expr: &Expr, key: &str) -> Result<String> {
344    match map_field(expr, key)? {
345        Expr::Symbol(symbol) => Ok(symbol.name.as_ref().to_owned()),
346        _ => Err(Error::Eval(format!(
347            "openai SSE field {key} must be a symbol"
348        ))),
349    }
350}
351
352fn symbol_value_field(expr: &Expr, key: &str) -> Result<Symbol> {
353    match map_field(expr, key)? {
354        Expr::Symbol(symbol) => Ok(symbol.clone()),
355        _ => Err(Error::Eval(format!(
356            "openai SSE field {key} must be a symbol"
357        ))),
358    }
359}
360
361fn map_field<'a>(expr: &'a Expr, key: &str) -> Result<&'a Expr> {
362    let Expr::Map(entries) = expr else {
363        return Err(Error::Eval("openai SSE payload must be a map".to_owned()));
364    };
365    sim_value::access::entry_field(entries, key)
366        .ok_or_else(|| Error::Eval(format!("openai SSE payload missing {key}")))
367}
368
369fn optional_u64_field(entries: &[(Expr, Expr)], key: &str) -> Result<Option<u64>> {
370    let Some(value) = entries.iter().find_map(|(field, value)| match field {
371        Expr::Symbol(symbol) if symbol.name.as_ref() == key => Some(value),
372        _ => None,
373    }) else {
374        return Ok(None);
375    };
376    match value {
377        Expr::Number(number) => number
378            .canonical
379            .parse::<u64>()
380            .map(Some)
381            .map_err(|err| Error::Eval(format!("openai SSE invalid {key}: {err}"))),
382        Expr::String(text) => text
383            .parse::<u64>()
384            .map(Some)
385            .map_err(|err| Error::Eval(format!("openai SSE invalid {key}: {err}"))),
386        _ => Err(Error::Eval(format!(
387            "openai SSE field {key} must be a number"
388        ))),
389    }
390}