use crate::provider::{LLMError, LLMStream, LLMStreamEvent};
pub(crate) struct TaskAbortGuard(pub(crate) tokio::task::JoinHandle<()>);
impl Drop for TaskAbortGuard {
fn drop(&mut self) {
self.0.abort();
}
}
pub(crate) fn spawn_openai_compatible_stream(
response: reqwest::Response,
provider_name: &'static str,
model: String,
reasoning_fields: &'static [&'static str],
delta_order: crate::providers::shared::OpenAiDeltaOrder,
include_cache_metrics: bool,
) -> LLMStream {
use async_stream::try_stream;
let bytes_stream = response.bytes_stream();
let (event_tx, event_rx) = tokio::sync::mpsc::unbounded_channel::<Result<LLMStreamEvent, LLMError>>();
let tx = event_tx.clone();
let stream_timeout = std::time::Duration::from_secs(300);
let handle = tokio::spawn(async move {
let aggregator_model = model.clone();
let mut aggregator = crate::providers::shared::StreamAggregator::new(aggregator_model);
let result = tokio::time::timeout(
stream_timeout,
crate::providers::shared::process_openai_stream(bytes_stream, provider_name, model, |value| {
crate::providers::shared::handle_openai_compatible_chunk(
&value,
&mut aggregator,
&tx,
reasoning_fields,
delta_order,
include_cache_metrics,
);
Ok(())
}),
)
.await;
match result {
Ok(Ok(_)) => {
let response = aggregator.finalize();
let _ = tx.send(Ok(LLMStreamEvent::Completed { response: Box::new(response) }));
}
Ok(Err(err)) => {
let _ = tx.send(Err(err));
}
Err(_elapsed) => {
let _ = tx.send(Err(LLMError::Provider {
message: format!("{provider_name}: streaming timed out after 5 minutes"),
metadata: None,
}));
}
}
});
let stream = try_stream! {
let mut receiver = event_rx;
let _abort_guard = TaskAbortGuard(handle);
while let Some(event) = receiver.recv().await {
yield event?;
}
};
Box::pin(stream)
}