use crate::brain::provider::OpenAIProvider;
use crate::brain::provider::Provider;
use crate::brain::provider::{ContentBlock, ContentDelta, LLMRequest, Message, StreamEvent};
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio::time::timeout;
fn chunk_role(id: &str) -> String {
format!(
r#"{{"id":"{id}","object":"chat.completion.chunk","model":"m","choices":[{{"index":0,"delta":{{"role":"assistant"}},"finish_reason":null}}]}}"#
)
}
fn chunk_text_empty_finish(id: &str) -> String {
format!(
r#"{{"id":"{id}","object":"chat.completion.chunk","model":"m","choices":[{{"index":0,"delta":{{"content":"Hi "}},"finish_reason":""}}]}}"#
)
}
fn chunk_tool_partial_empty_finish(id: &str) -> String {
format!(
r#"{{"id":"{id}","object":"chat.completion.chunk","model":"m","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":0,"id":"call_1","type":"function","function":{{"name":"get_weather","arguments":"{{\"city\":"}}}}]}},"finish_reason":""}}]}}"#
)
}
fn chunk_tool_rest(id: &str) -> String {
format!(
r#"{{"id":"{id}","object":"chat.completion.chunk","model":"m","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":0,"function":{{"arguments":"\"Paris\"}}"}}}}]}},"finish_reason":null}}]}}"#
)
}
fn chunk_terminal(id: &str) -> String {
format!(
r#"{{"id":"{id}","object":"chat.completion.chunk","model":"m","choices":[{{"index":0,"delta":{{}},"finish_reason":"tool_calls"}}]}}"#
)
}
async fn serve_sse(listener: TcpListener, body: String) {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut buf = [0u8; 8192];
let _ = timeout(Duration::from_secs(5), sock.read(&mut buf)).await;
let resp = format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
body.len(),
body
);
sock.write_all(resp.as_bytes()).await.expect("write sse");
sock.flush().await.ok();
}
async fn collect_events(provider: &OpenAIProvider) -> Vec<StreamEvent> {
let req = LLMRequest::new("test-model", vec![Message::user("weather?")]);
let mut stream = provider.stream(req).await.expect("stream opens");
let mut events = Vec::new();
while let Some(ev) = futures::StreamExt::next(&mut stream).await {
let ev = ev.expect("event ok");
let done = matches!(ev, StreamEvent::MessageStop);
events.push(ev);
if done {
break;
}
}
events
}
#[tokio::test]
async fn empty_finish_reason_does_not_flush_partial_tool_calls() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let port = listener.local_addr().unwrap().port();
let id = "chatcmpl-105";
let sse = format!(
"data: {}\n\ndata: {}\n\ndata: {}\n\ndata: {}\n\ndata: {}\n\ndata: [DONE]\n\n",
chunk_role(id),
chunk_text_empty_finish(id),
chunk_tool_partial_empty_finish(id),
chunk_tool_rest(id),
chunk_terminal(id),
);
tokio::spawn(serve_sse(listener, sse));
let provider = OpenAIProvider::local(format!("http://127.0.0.1:{port}/chat/completions"));
let events = timeout(Duration::from_secs(10), collect_events(&provider))
.await
.expect("stream completes in time");
let tool_starts: Vec<&ContentBlock> = events
.iter()
.filter_map(|e| match e {
StreamEvent::ContentBlockStart {
content_block: block @ ContentBlock::ToolUse { .. },
..
} => Some(block),
_ => None,
})
.collect();
assert_eq!(tool_starts.len(), 1, "tool starts: {tool_starts:?}");
match &tool_starts[0] {
ContentBlock::ToolUse { name, input, .. } => {
assert_eq!(name, "get_weather");
assert_eq!(input["city"], "Paris", "args must be complete: {input}");
}
other => panic!("expected ToolUse, got {other:?}"),
}
let deltas: Vec<_> = events
.iter()
.filter_map(|e| match e {
StreamEvent::MessageDelta { delta, .. } => Some(delta.stop_reason.clone()),
_ => None,
})
.collect();
assert_eq!(deltas.len(), 1, "message deltas: {deltas:?}");
assert!(deltas[0].is_some(), "terminal delta carries a stop reason");
let text: String = events
.iter()
.filter_map(|e| match e {
StreamEvent::ContentBlockDelta {
delta: ContentDelta::TextDelta { text },
..
} => Some(text.clone()),
_ => None,
})
.collect();
assert_eq!(text, "Hi ");
assert!(matches!(events.last(), Some(StreamEvent::MessageStop)));
}