use lc_core::tools::ToolCall;
use serde::Deserialize;
pub struct SseByteFramer {
buf: Vec<u8>,
}
impl SseByteFramer {
pub fn new() -> Self {
Self { buf: Vec::new() }
}
pub fn push(&mut self, chunk: &[u8]) -> Vec<String> {
self.buf.extend_from_slice(chunk);
let mut events = Vec::new();
while let Some((sep_len, end)) = find_event_end(&self.buf) {
let raw: Vec<u8> = self.buf.drain(..end).collect();
self.buf.drain(..sep_len);
events.push(format!("{}\n\n", String::from_utf8_lossy(&raw)));
}
events
}
pub fn pending(&self) -> usize {
self.buf.len()
}
}
impl Default for SseByteFramer {
fn default() -> Self {
Self::new()
}
}
fn find_event_end(buf: &[u8]) -> Option<(usize, usize)> {
let mut i = 0;
while i < buf.len() {
if buf[i] == b'\n' {
if i + 1 < buf.len() && buf[i + 1] == b'\n' {
return Some((2, i));
}
if i + 2 < buf.len()
&& buf[i + 1] == b'\r'
&& buf[i + 2] == b'\n'
&& i >= 1
&& buf[i - 1] == b'\r'
{
return Some((4, i - 1));
}
}
i += 1;
}
None
}
pub struct SSEParser {
buffer: String,
}
impl SSEParser {
pub fn new() -> Self {
Self {
buffer: String::new(),
}
}
pub fn parse(&mut self, chunk: &str) -> Vec<SSEEvent> {
self.buffer.push_str(&chunk.replace("\r\n", "\n"));
let mut events = Vec::new();
while let Some(pos) = self.buffer.find("\n\n") {
let event_text = self.buffer[..pos].to_string();
self.buffer.drain(..=pos + 1);
if let Some(event) = self.parse_event(&event_text) {
events.push(event);
}
}
events
}
fn parse_event(&self, text: &str) -> Option<SSEEvent> {
let mut event_type = None;
let mut data_lines: Vec<String> = Vec::new();
for line in text.lines() {
if let Some(value) = line.strip_prefix("event:") {
event_type = Some(value.trim().to_string());
} else if let Some(value) = line.strip_prefix("data:") {
data_lines.push(value.trim().to_string());
}
}
if data_lines.is_empty() {
None
} else {
Some(SSEEvent {
event: event_type,
data: data_lines.join("\n"),
})
}
}
}
impl Default for SSEParser {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct SSEEvent {
pub event: Option<String>,
pub data: String,
}
impl SSEEvent {
pub fn is_done(&self) -> bool {
self.data == "[DONE]"
}
pub fn parse_openai_chunk(&self) -> Result<Option<OpenAIStreamChunk>, serde_json::Error> {
if self.is_done() {
return Ok(None);
}
let chunk: OpenAIStreamChunk = serde_json::from_str(&self.data)?;
Ok(Some(chunk))
}
}
#[derive(Debug, Deserialize)]
pub struct OpenAIStreamChunk {
pub id: String,
pub object: String,
pub created: i64,
pub model: String,
pub choices: Vec<StreamChoice>,
#[serde(default)]
pub usage: Option<StreamUsage>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct StreamUsage {
pub prompt_tokens: usize,
pub completion_tokens: usize,
pub total_tokens: usize,
}
#[derive(Debug, Deserialize)]
pub struct StreamChoice {
pub index: i32,
pub delta: Delta,
#[serde(default)]
pub finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct Delta {
#[serde(default)]
pub role: Option<String>,
#[serde(default)]
pub content: Option<String>,
#[serde(default)]
pub tool_calls: Option<Vec<StreamToolCallDelta>>,
}
#[derive(Debug, Deserialize)]
pub struct StreamToolCallDelta {
#[serde(default)]
pub index: usize,
#[serde(default)]
pub id: Option<String>,
#[serde(default, rename = "type")]
pub tool_type: Option<String>,
#[serde(default)]
pub function: Option<StreamFunctionDelta>,
}
#[derive(Debug, Deserialize)]
pub struct StreamFunctionDelta {
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub arguments: Option<String>,
}
#[derive(Default)]
pub struct StreamToolCallAccumulator {
calls: Vec<(usize, AccumulatedToolCall)>,
}
#[derive(Default)]
struct AccumulatedToolCall {
id: String,
name: String,
arguments: String,
}
impl StreamToolCallAccumulator {
pub fn push(&mut self, delta: &StreamToolCallDelta) {
let slot = match self.calls.iter_mut().find(|(i, _)| *i == delta.index) {
Some(slot) => slot,
None => {
self.calls
.push((delta.index, AccumulatedToolCall::default()));
self.calls.last_mut().expect("just pushed")
}
};
if let Some(id) = &delta.id {
if slot.1.id.is_empty() {
slot.1.id.clone_from(id);
}
}
if let Some(function) = &delta.function {
if let Some(name) = &function.name {
if slot.1.name.is_empty() {
slot.1.name.clone_from(name);
}
}
if let Some(arguments) = &function.arguments {
slot.1.arguments.push_str(arguments);
}
}
}
pub fn build(&self) -> Vec<ToolCall> {
self.calls
.iter()
.filter(|(_, c)| !c.name.is_empty())
.map(|(_, c)| {
ToolCall::builder(&c.id)
.name(&c.name)
.arguments(&c.arguments)
.build()
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sse_parser() {
let mut parser = SSEParser::new();
let chunk = "data: {\"test\": \"value\"}\n\n";
let events = parser.parse(chunk);
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "{\"test\": \"value\"}");
}
#[test]
fn test_sse_done_event() {
let mut parser = SSEParser::new();
let chunk = "data: [DONE]\n\n";
let events = parser.parse(chunk);
assert_eq!(events.len(), 1);
assert!(events[0].is_done());
}
#[test]
fn test_sse_parser_crlf_event_separator() {
let mut parser = SSEParser::new();
let events = parser.parse("data: {\"a\":1}\r\n\r\n");
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "{\"a\":1}");
}
#[test]
fn test_sse_parser_crlf_multiple_events_across_chunks() {
let mut parser = SSEParser::new();
let events = parser.parse("data: {\"a\":1}\r\n\r\ndata: {\"b\":2}\r\n\r\n");
assert_eq!(events.len(), 2);
assert_eq!(events[0].data, "{\"a\":1}");
assert_eq!(events[1].data, "{\"b\":2}");
}
#[test]
fn test_sse_parser_crlf_split_across_chunks() {
let mut parser = SSEParser::new();
assert_eq!(parser.parse("data: {\"a\":1}\r\n").len(), 0);
let events = parser.parse("\r\n");
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "{\"a\":1}");
}
#[test]
fn test_sse_parser_multiline_data_joined_with_newline() {
let mut parser = SSEParser::new();
let events = parser.parse("data: line1\ndata: line2\ndata: line3\n\n");
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "line1\nline2\nline3");
}
#[test]
fn test_sse_parser_data_without_space() {
let mut parser = SSEParser::new();
let events = parser.parse("data:{\"no\":\"space\"}\n\n");
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "{\"no\":\"space\"}");
}
#[test]
fn test_openai_chunk_parsing() {
let event = SSEEvent {
event: None,
data: r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-3.5-turbo","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}"#.to_string(),
};
let chunk = event.parse_openai_chunk().unwrap().unwrap();
assert_eq!(chunk.choices[0].delta.content, Some("Hello".to_string()));
}
#[test]
fn test_openai_chunk_parsing_tool_calls_delta() {
let event = SSEEvent {
event: None,
data: r#"{"id":"chatcmpl-1","object":"chat.completion.chunk","created":1,"model":"gpt","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"get_weather","arguments":""}}]},"finish_reason":null}]}"#.to_string(),
};
let chunk = event.parse_openai_chunk().unwrap().unwrap();
let deltas = chunk.choices[0]
.delta
.tool_calls
.as_ref()
.expect("tool_calls parsed");
assert_eq!(deltas.len(), 1);
assert_eq!(deltas[0].index, 0);
assert_eq!(deltas[0].id.as_deref(), Some("call_1"));
assert_eq!(
deltas[0].tool_type.as_deref(),
Some("function"),
"serde(rename=\"type\") maps the wire field"
);
assert_eq!(
deltas[0].function.as_ref().and_then(|f| f.name.as_deref()),
Some("get_weather")
);
}
#[test]
fn test_stream_tool_call_accumulator_reconstructs_fragmented_calls() {
let mut acc = StreamToolCallAccumulator::default();
acc.push(
&serde_json::from_str::<StreamToolCallDelta>(
r#"{"index":0,"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"city\":\"beij"}}"#,
)
.unwrap(),
);
acc.push(
&serde_json::from_str::<StreamToolCallDelta>(
r#"{"index":0,"function":{"arguments":"ing\"}"}}"#,
)
.unwrap(),
);
acc.push(
&serde_json::from_str::<StreamToolCallDelta>(
r#"{"index":1,"id":"call_2","type":"function","function":{"name":"add","arguments":"{\"a\":1}"}}"#,
)
.unwrap(),
);
let calls = acc.build();
assert_eq!(calls.len(), 2, "two distinct tool calls by index");
assert_eq!(calls[0].id, "call_1");
assert_eq!(calls[0].name(), "get_weather");
assert_eq!(calls[0].arguments(), r#"{"city":"beijing"}"#);
assert_eq!(calls[1].id, "call_2");
assert_eq!(calls[1].name(), "add");
assert_eq!(calls[1].arguments(), r#"{"a":1}"#);
}
#[test]
fn test_stream_tool_call_accumulator_drops_unfinished_call() {
let mut acc = StreamToolCallAccumulator::default();
acc.push(
&serde_json::from_str::<StreamToolCallDelta>(
r#"{"index":0,"id":"call_1","function":{"name":"get_weather"}}"#,
)
.unwrap(),
);
acc.push(
&serde_json::from_str::<StreamToolCallDelta>(
r#"{"index":1,"function":{"arguments":"{\"x\":1}"}}"#,
)
.unwrap(),
);
let calls = acc.build();
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].name(), "get_weather");
assert_eq!(calls[0].arguments(), "");
}
}