use std::sync::Arc;
use tokio::sync::mpsc;
use crate::models::{ReasoningChunk, StreamCallback, StreamEvent as ModelStreamEvent};
use super::super::ctx::StreamEvent;
pub fn ordered_relay(
bounded_sink: mpsc::Sender<StreamEvent>,
) -> (
mpsc::UnboundedSender<StreamEvent>,
crate::utils::AbortOnDrop,
) {
let (tx, mut rx) = mpsc::unbounded_channel::<StreamEvent>();
let handle = crate::utils::spawn_guarded(async move {
while let Some(event) = rx.recv().await {
if bounded_sink.send(event).await.is_err() {
break;
}
}
});
(tx, handle)
}
pub fn forward_callback(sink: mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
Arc::new(move |event: ModelStreamEvent| {
let mapped = match event {
ModelStreamEvent::Text(s) => StreamEvent::Text(s),
ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
text: chunk.text,
signature: chunk.signature,
}),
ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
ModelStreamEvent::Status(s) => StreamEvent::Status(s),
ModelStreamEvent::Done { .. } => StreamEvent::Done {
usage: None,
provider_continuation: None,
stop_reason: None,
},
};
let _ = sink.send(mapped);
})
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn events_arrive_in_order() {
let (sink_tx, mut sink_rx) = mpsc::channel::<StreamEvent>(16);
let (relay, _handle) = ordered_relay(sink_tx);
relay.send(StreamEvent::Text("a".to_string())).unwrap();
relay
.send(StreamEvent::Reasoning(ReasoningChunk {
text: "r1".to_string(),
signature: None,
}))
.unwrap();
relay.send(StreamEvent::Text("b".to_string())).unwrap();
relay
.send(StreamEvent::Done {
usage: None,
provider_continuation: None,
stop_reason: None,
})
.unwrap();
drop(relay);
let mut seen: Vec<&'static str> = Vec::new();
while let Some(ev) = sink_rx.recv().await {
seen.push(match ev {
StreamEvent::Text(s) if s == "a" => "text-a",
StreamEvent::Text(s) if s == "b" => "text-b",
StreamEvent::Reasoning(_) => "reasoning",
StreamEvent::Done { .. } => "done",
_ => "other",
});
}
assert_eq!(seen, vec!["text-a", "reasoning", "text-b", "done"]);
}
#[tokio::test]
async fn stream_callback_forwards_text_event() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
let cb = forward_callback(tx);
cb(ModelStreamEvent::Text("hello".to_string()));
let recv = tokio::time::timeout(std::time::Duration::from_millis(100), rx.recv())
.await
.expect("recv")
.expect("sender alive");
match recv {
StreamEvent::Text(s) => assert_eq!(s, "hello"),
_ => panic!("wrong variant"),
}
}
#[tokio::test]
async fn stream_callback_forwards_status_notice() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
let cb = forward_callback(tx);
cb(ModelStreamEvent::Status(
"Starting the local Ollama server…".to_string(),
));
let recv = tokio::time::timeout(std::time::Duration::from_millis(100), rx.recv())
.await
.expect("recv")
.expect("sender alive");
match recv {
StreamEvent::Status(s) => assert!(s.contains("Starting")),
_ => panic!("wrong variant"),
}
}
#[tokio::test]
async fn stream_callback_done_never_invents_usage() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
let cb = forward_callback(tx);
cb(ModelStreamEvent::Done { tokens: 42 });
let recv = tokio::time::timeout(std::time::Duration::from_millis(100), rx.recv())
.await
.expect("recv")
.expect("sender");
match recv {
StreamEvent::Done { usage, .. } => assert!(usage.is_none()),
_ => panic!("wrong variant"),
}
}
}