use futures_util::StreamExt;
use saya_agent::{
CancellationToken, ChatProvider, ChatRequest, MAX_STREAM_BYTES, ProviderError, ProviderEvent,
TokenUsage,
};
#[derive(Debug)]
pub(crate) struct ExtractionReply {
pub content: String,
pub usage: Option<TokenUsage>,
}
#[derive(Debug, thiserror::Error)]
pub(crate) enum ExtractionStreamError {
#[error("extraction reply is not JSON")]
NotJson { usage: Option<TokenUsage> },
#[error(transparent)]
Provider(ProviderError),
}
pub(crate) async fn collect_extraction(
provider: &dyn ChatProvider,
request: ChatRequest,
) -> Result<ExtractionReply, ExtractionStreamError> {
let cancellation = CancellationToken::new();
let mut stream = provider
.stream(request, cancellation.clone())
.await
.map_err(ExtractionStreamError::Provider)?;
let (mut content, mut usage) = (String::new(), None);
let (mut bytes, mut decided) = (0usize, false);
while let Some(event) = stream.next().await {
match event.map_err(ExtractionStreamError::Provider)? {
ProviderEvent::TextDelta(delta) => {
bytes = bounded(bytes, &delta)?;
content.push_str(&delta);
if !decided && let Some(first) = content.chars().find(|c| !c.is_whitespace()) {
decided = true;
if first != '{' && first != '`' {
cancellation.cancel();
return Err(ExtractionStreamError::NotJson { usage });
}
}
}
ProviderEvent::ReasoningDelta(delta) => bytes = bounded(bytes, &delta)?,
ProviderEvent::Usage(reported) => usage = Some(reported),
ProviderEvent::ToolCalls(_) => {}
ProviderEvent::Done => break,
_ => {}
}
}
Ok(ExtractionReply { content, usage })
}
fn bounded(accumulated: usize, delta: &str) -> Result<usize, ExtractionStreamError> {
let total = accumulated.saturating_add(delta.len());
if total > MAX_STREAM_BYTES {
return Err(ExtractionStreamError::Provider(ProviderError::Request(
"provider stream exceeded size limit".into(),
)));
}
Ok(total)
}
#[cfg(test)]
#[path = "extraction_stream_tests.rs"]
mod tests;