Skip to main content

starweaver_model/transport/
client.rs

1use std::{collections::VecDeque, sync::Arc};
2
3use async_trait::async_trait;
4use serde_json::Value;
5use starweaver_core::CancellationToken;
6
7use crate::ModelError;
8
9use super::{HttpRequest, HttpResponse};
10
11/// Async HTTP client abstraction used by production model adapters.
12#[async_trait]
13pub trait ModelHttpClient: Send + Sync {
14    /// Send a JSON model request.
15    ///
16    /// # Errors
17    ///
18    /// Returns an error when transport, status, or response decoding fails.
19    async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError>;
20
21    /// Send a server-sent events model request and return JSON `data:` payloads.
22    ///
23    /// # Errors
24    ///
25    /// Returns an error when transport, status, or event decoding fails.
26    async fn send_event_stream(&self, request: HttpRequest) -> Result<Vec<Value>, ModelError> {
27        let mut stream = self.send_event_stream_incremental(request).await?;
28        let mut events = Vec::new();
29        while let Some(event) = stream.recv().await {
30            events.push(event?);
31        }
32        Ok(events)
33    }
34
35    /// Send a server-sent events model request and return events as they arrive.
36    ///
37    /// # Errors
38    ///
39    /// Returns an error when transport setup fails.
40    async fn send_event_stream_incremental(
41        &self,
42        request: HttpRequest,
43    ) -> Result<ModelEventStream, ModelError> {
44        Err(ModelError::Transport(format!(
45            "server-sent event streaming is not implemented for {}",
46            request.url
47        )))
48    }
49
50    /// Send a WebSocket model request and return JSON text-frame events as they arrive.
51    ///
52    /// # Errors
53    ///
54    /// Returns an error when WebSocket transport setup fails.
55    async fn send_websocket_event_stream_incremental(
56        &self,
57        request: HttpRequest,
58    ) -> Result<ModelEventStream, ModelError> {
59        Err(ModelError::Transport(format!(
60            "websocket event streaming is not implemented for {}",
61            request.url
62        )))
63    }
64
65    /// Create a session that may reuse WebSocket transport state across sequential requests.
66    fn websocket_event_session(&self) -> Box<dyn ModelWebSocketEventSession + '_> {
67        Box::new(PerRequestWebSocketEventSession { client: self })
68    }
69}
70
71/// Run-scoped WebSocket event transport session.
72#[async_trait]
73pub trait ModelWebSocketEventSession: Send {
74    /// Send a WebSocket model request and return JSON text-frame events as they arrive.
75    async fn send_websocket_event_stream_incremental(
76        &mut self,
77        request: HttpRequest,
78    ) -> Result<ModelEventStream, ModelError>;
79
80    /// Reset any reusable transport state held by the session.
81    async fn reset(&mut self) {}
82}
83
84struct PerRequestWebSocketEventSession<'a, C: ModelHttpClient + ?Sized> {
85    client: &'a C,
86}
87
88#[async_trait]
89impl<C> ModelWebSocketEventSession for PerRequestWebSocketEventSession<'_, C>
90where
91    C: ModelHttpClient + ?Sized,
92{
93    async fn send_websocket_event_stream_incremental(
94        &mut self,
95        request: HttpRequest,
96    ) -> Result<ModelEventStream, ModelError> {
97        self.client
98            .send_websocket_event_stream_incremental(request)
99            .await
100    }
101}
102
103/// Receiver for incremental model JSON events.
104pub struct ModelEventStream {
105    receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>,
106    prefetched: VecDeque<Result<Value, ModelError>>,
107    cancellation_token: CancellationToken,
108    drop_abort_token: Option<CancellationToken>,
109}
110
111impl ModelEventStream {
112    /// Build a stream from a channel receiver.
113    #[must_use]
114    pub fn new(receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>) -> Self {
115        Self::new_with_cancellation(receiver, CancellationToken::default())
116    }
117
118    /// Build a stream from a channel receiver and cancellation token.
119    #[must_use]
120    pub const fn new_with_cancellation(
121        receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>,
122        cancellation_token: CancellationToken,
123    ) -> Self {
124        Self::new_with_cancellation_and_drop_abort(receiver, cancellation_token, None)
125    }
126
127    /// Build a stream from a channel receiver, cancellation token, and drop-abort token.
128    #[must_use]
129    pub const fn new_with_cancellation_and_drop_abort(
130        receiver: tokio::sync::mpsc::Receiver<Result<Value, ModelError>>,
131        cancellation_token: CancellationToken,
132        drop_abort_token: Option<CancellationToken>,
133    ) -> Self {
134        Self {
135            receiver,
136            prefetched: VecDeque::new(),
137            cancellation_token,
138            drop_abort_token,
139        }
140    }
141
142    pub(crate) fn prepend_events(
143        mut self,
144        events: impl IntoIterator<Item = Result<Value, ModelError>>,
145    ) -> Self {
146        let mut prefetched = events.into_iter().collect::<VecDeque<_>>();
147        prefetched.append(&mut self.prefetched);
148        self.prefetched = prefetched;
149        self
150    }
151
152    /// Return the transport-local drop-abort token, when the stream owns one.
153    #[must_use]
154    pub fn drop_abort_token(&self) -> Option<CancellationToken> {
155        self.drop_abort_token.clone()
156    }
157
158    /// Receive the next JSON event from the stream.
159    pub async fn recv(&mut self) -> Option<Result<Value, ModelError>> {
160        if self.cancellation_token.is_cancelled() {
161            return Some(Err(ModelError::Cancelled {
162                reason: "model event stream cancellation requested".to_string(),
163            }));
164        }
165        if let Some(event) = self.prefetched.pop_front() {
166            return Some(event);
167        }
168        tokio::select! {
169            biased;
170            () = self.cancellation_token.cancelled() => Some(Err(ModelError::Cancelled {
171                reason: "model event stream cancellation requested".to_string(),
172            })),
173            event = self.receiver.recv() => event,
174        }
175    }
176}
177
178impl Drop for ModelEventStream {
179    fn drop(&mut self) {
180        if let Some(token) = &self.drop_abort_token {
181            token.cancel();
182        }
183    }
184}
185
186/// Shared reference to an HTTP client.
187pub type DynHttpClient = Arc<dyn ModelHttpClient>;