starweaver_model/transport/
client.rs1use 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_trait]
13pub trait ModelHttpClient: Send + Sync {
14 async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError>;
20
21 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 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 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 fn websocket_event_session(&self) -> Box<dyn ModelWebSocketEventSession + '_> {
67 Box::new(PerRequestWebSocketEventSession { client: self })
68 }
69}
70
71#[async_trait]
73pub trait ModelWebSocketEventSession: Send {
74 async fn send_websocket_event_stream_incremental(
76 &mut self,
77 request: HttpRequest,
78 ) -> Result<ModelEventStream, ModelError>;
79
80 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
103pub 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 #[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 #[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 #[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 #[must_use]
154 pub fn drop_abort_token(&self) -> Option<CancellationToken> {
155 self.drop_abort_token.clone()
156 }
157
158 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
186pub type DynHttpClient = Arc<dyn ModelHttpClient>;