use std::sync::Arc;
use async_trait::async_trait;
use serde_json::Value;
use starweaver_core::CancellationToken;
use crate::ModelError;
use super::{HttpRequest, HttpResponse};
#[async_trait]
pub trait ModelHttpClient: Send + Sync {
async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError>;
async fn send_event_stream(&self, request: HttpRequest) -> Result<Vec<Value>, ModelError> {
let mut stream = self.send_event_stream_incremental(request).await?;
let mut events = Vec::new();
while let Some(event) = stream.recv().await {
events.push(event?);
}
Ok(events)
}
async fn send_event_stream_incremental(
&self,
request: HttpRequest,
) -> Result<ModelEventStream, ModelError> {
Err(ModelError::Transport(format!(
"server-sent event streaming is not implemented for {}",
request.url
)))
}
}
pub struct ModelEventStream {
receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>,
cancellation_token: CancellationToken,
}
impl ModelEventStream {
#[must_use]
pub fn new(receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>) -> Self {
Self::new_with_cancellation(receiver, CancellationToken::default())
}
#[must_use]
pub const fn new_with_cancellation(
receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>,
cancellation_token: CancellationToken,
) -> Self {
Self {
receiver,
cancellation_token,
}
}
pub async fn recv(&mut self) -> Option<Result<Value, ModelError>> {
if self.cancellation_token.is_cancelled() {
return Some(Err(ModelError::Cancelled {
reason: "model event stream cancellation requested".to_string(),
}));
}
tokio::select! {
biased;
() = self.cancellation_token.cancelled() => Some(Err(ModelError::Cancelled {
reason: "model event stream cancellation requested".to_string(),
})),
event = self.receiver.recv() => event,
}
}
}
pub type DynHttpClient = Arc<dyn ModelHttpClient>;