1use std::{any::Any, cell::Cell, fmt, marker::PhantomData, rc::Rc};
2
3use futures::{
4 SinkExt, StreamExt,
5 channel::mpsc,
6 future::{Either, LocalBoxFuture, select},
7 lock::Mutex,
8};
9use lenso_kernel::{
10 CancellationToken, InvocationContext, NativeStreamItem, NativeStreamSession, RuntimeFailure,
11 StreamCapability,
12};
13
14use crate::PluginResult;
15
16#[derive(Debug)]
18pub enum StreamInput<C: StreamCapability> {
19 Message(C::Message),
21 PeerHalfClosed,
23}
24
25enum ProviderOutput<C: StreamCapability> {
26 Message(C::Message),
27 PeerHalfClosed,
28 Terminal(Result<(), C::DomainError>),
29 Runtime(RuntimeFailure),
30}
31
32pub struct ProviderStreamChannel<C: StreamCapability> {
37 outgoing: mpsc::Sender<ProviderOutput<C>>,
38 incoming: mpsc::Receiver<StreamInput<C>>,
39 cancellation: CancellationToken,
40 send_closed: bool,
41 terminated: bool,
42}
43
44impl<C: StreamCapability> fmt::Debug for ProviderStreamChannel<C> {
45 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
46 formatter
47 .debug_struct("ProviderStreamChannel")
48 .field("capability", &C::ID)
49 .field("send_closed", &self.send_closed)
50 .field("terminated", &self.terminated)
51 .field("cancelled", &self.cancellation.is_cancelled())
52 .finish_non_exhaustive()
53 }
54}
55
56impl<C: StreamCapability> ProviderStreamChannel<C> {
57 pub async fn send(&mut self, message: C::Message) -> Result<(), RuntimeFailure> {
59 if self.send_closed || self.terminated {
60 return Err(RuntimeFailure::AdmissionClosed);
61 }
62 self.send_output(ProviderOutput::Message(message)).await
63 }
64
65 pub async fn close_send(&mut self) -> Result<(), RuntimeFailure> {
67 if self.send_closed || self.terminated {
68 return Err(RuntimeFailure::AdmissionClosed);
69 }
70 self.send_closed = true;
71 self.send_output(ProviderOutput::PeerHalfClosed).await
72 }
73
74 pub async fn finish(&mut self) -> Result<(), RuntimeFailure> {
76 self.terminate(ProviderOutput::Terminal(Ok(()))).await
77 }
78
79 pub async fn fail(&mut self, error: C::DomainError) -> Result<(), RuntimeFailure> {
81 self.terminate(ProviderOutput::Terminal(Err(error))).await
82 }
83
84 pub async fn fail_runtime(&mut self, error: RuntimeFailure) -> Result<(), RuntimeFailure> {
86 self.terminate(ProviderOutput::Runtime(error)).await
87 }
88
89 pub async fn complete(
94 mut self,
95 result: PluginResult<(), C::DomainError>,
96 ) -> Result<(), RuntimeFailure> {
97 if !self.send_closed {
98 self.close_send().await?;
99 }
100 match result {
101 Ok(()) => self.finish().await,
102 Err(crate::PluginError::Domain(error)) => self.fail(error).await,
103 Err(crate::PluginError::Runtime(error)) => self.fail_runtime(error).await,
104 }
105 }
106
107 pub async fn receive(&mut self) -> Result<StreamInput<C>, RuntimeFailure> {
109 if self.cancellation.is_cancelled() {
110 return Err(RuntimeFailure::AdmissionClosed);
111 }
112 let receive = self.incoming.next();
113 futures::pin_mut!(receive);
114 match select(receive, self.cancellation.cancelled()).await {
115 Either::Left((Some(input), _)) => Ok(input),
116 Either::Left((None, _)) | Either::Right(_) => Err(RuntimeFailure::AdmissionClosed),
117 }
118 }
119
120 pub fn is_cancelled(&self) -> bool {
122 self.cancellation.is_cancelled()
123 }
124
125 async fn terminate(&mut self, output: ProviderOutput<C>) -> Result<(), RuntimeFailure> {
126 if self.terminated {
127 return Err(RuntimeFailure::AdmissionClosed);
128 }
129 self.terminated = true;
130 self.send_output(output).await
131 }
132
133 async fn send_output(&mut self, output: ProviderOutput<C>) -> Result<(), RuntimeFailure> {
134 if self.cancellation.is_cancelled() {
135 return Err(RuntimeFailure::AdmissionClosed);
136 }
137 let send = self.outgoing.send(output);
138 futures::pin_mut!(send);
139 match select(send, self.cancellation.cancelled()).await {
140 Either::Left((Ok(()), _)) => Ok(()),
141 Either::Left((Err(_), _)) | Either::Right(_) => Err(RuntimeFailure::AdmissionClosed),
142 }
143 }
144}
145
146pub struct ProviderStream<C: StreamCapability> {
148 incoming: mpsc::Sender<StreamInput<C>>,
149 outgoing: Rc<Mutex<mpsc::Receiver<ProviderOutput<C>>>>,
150 cancellation: CancellationToken,
151 consumer_send_closed: Rc<Cell<bool>>,
152 terminated: Rc<Cell<bool>>,
153 marker: PhantomData<fn() -> C>,
154}
155
156impl<C: StreamCapability> ProviderStream<C> {
157 pub fn channel(
159 context: &InvocationContext,
160 capacity: usize,
161 ) -> (Self, ProviderStreamChannel<C>) {
162 let (incoming_sender, incoming_receiver) = mpsc::channel(capacity);
163 let (outgoing_sender, outgoing_receiver) = mpsc::channel(capacity);
164 let cancellation = context.cancellation();
165 (
166 Self {
167 incoming: incoming_sender,
168 outgoing: Rc::new(Mutex::new(outgoing_receiver)),
169 cancellation: cancellation.clone(),
170 consumer_send_closed: Rc::new(Cell::new(false)),
171 terminated: Rc::new(Cell::new(false)),
172 marker: PhantomData,
173 },
174 ProviderStreamChannel {
175 outgoing: outgoing_sender,
176 incoming: incoming_receiver,
177 cancellation,
178 send_closed: false,
179 terminated: false,
180 },
181 )
182 }
183}
184
185impl<C: StreamCapability> fmt::Debug for ProviderStream<C> {
186 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
187 formatter
188 .debug_struct("ProviderStream")
189 .field("capability", &C::ID)
190 .field("consumer_send_closed", &self.consumer_send_closed.get())
191 .field("terminated", &self.terminated.get())
192 .field("cancelled", &self.cancellation.is_cancelled())
193 .finish_non_exhaustive()
194 }
195}
196
197impl<C: StreamCapability> NativeStreamSession for ProviderStream<C> {
198 fn send(&self, message: Box<dyn Any>) -> LocalBoxFuture<'static, Result<(), RuntimeFailure>> {
199 if self.cancellation.is_cancelled() || self.terminated.get() {
200 return Box::pin(futures::future::ready(Err(RuntimeFailure::AdmissionClosed)));
201 }
202 if self.consumer_send_closed.get() {
203 return Box::pin(futures::future::ready(Err(
204 RuntimeFailure::ProtocolViolation { capability: C::ID },
205 )));
206 }
207 let Ok(message) = message.downcast::<C::Message>() else {
208 return Box::pin(futures::future::ready(Err(
209 RuntimeFailure::ProtocolViolation { capability: C::ID },
210 )));
211 };
212 let mut incoming = self.incoming.clone();
213 let cancellation = self.cancellation.clone();
214 Box::pin(async move {
215 let send = incoming.send(StreamInput::Message(*message));
216 futures::pin_mut!(send);
217 match select(send, cancellation.cancelled()).await {
218 Either::Left((Ok(()), _)) => Ok(()),
219 Either::Left((Err(_), _)) | Either::Right(_) => {
220 Err(RuntimeFailure::AdmissionClosed)
221 }
222 }
223 })
224 }
225
226 fn receive(&self) -> LocalBoxFuture<'static, Result<NativeStreamItem, RuntimeFailure>> {
227 if self.cancellation.is_cancelled() || self.terminated.get() {
228 return Box::pin(futures::future::ready(Err(RuntimeFailure::AdmissionClosed)));
229 }
230 let outgoing = Rc::clone(&self.outgoing);
231 let cancellation = self.cancellation.clone();
232 let terminated = Rc::clone(&self.terminated);
233 Box::pin(async move {
234 let receive = async move { outgoing.lock().await.next().await };
235 futures::pin_mut!(receive);
236 match select(receive, cancellation.cancelled()).await {
237 Either::Left((Some(ProviderOutput::Message(message)), _)) => {
238 Ok(NativeStreamItem::Message(Box::new(message)))
239 }
240 Either::Left((Some(ProviderOutput::PeerHalfClosed), _)) => {
241 Ok(NativeStreamItem::PeerHalfClosed)
242 }
243 Either::Left((Some(ProviderOutput::Terminal(result)), _)) => {
244 terminated.set(true);
245 Ok(NativeStreamItem::Terminal(
246 result.map_err(|error| Box::new(error) as Box<dyn Any>),
247 ))
248 }
249 Either::Left((Some(ProviderOutput::Runtime(error)), _)) => {
250 terminated.set(true);
251 Err(error)
252 }
253 Either::Left((None, _)) => Err(RuntimeFailure::PluginFailure {
254 detail: format!("provider Stream {} ended without a terminal outcome", C::ID),
255 }),
256 Either::Right(_) => Err(RuntimeFailure::AdmissionClosed),
257 }
258 })
259 }
260
261 fn close_send(&self) -> LocalBoxFuture<'static, Result<(), RuntimeFailure>> {
262 if self.cancellation.is_cancelled() || self.terminated.get() {
263 return Box::pin(futures::future::ready(Err(RuntimeFailure::AdmissionClosed)));
264 }
265 if self.consumer_send_closed.replace(true) {
266 return Box::pin(futures::future::ready(Err(
267 RuntimeFailure::ProtocolViolation { capability: C::ID },
268 )));
269 }
270 let mut incoming = self.incoming.clone();
271 let cancellation = self.cancellation.clone();
272 Box::pin(async move {
273 let send = incoming.send(StreamInput::PeerHalfClosed);
274 futures::pin_mut!(send);
275 match select(send, cancellation.cancelled()).await {
276 Either::Left((Ok(()), _)) => Ok(()),
277 Either::Left((Err(_), _)) | Either::Right(_) => {
278 Err(RuntimeFailure::AdmissionClosed)
279 }
280 }
281 })
282 }
283
284 fn cancel(&self) {
285 self.cancellation.cancel();
286 }
287}