Skip to main content

lenso/
provider_stream.rs

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/// One typed value sent by a Capability consumer to a Stream provider.
17#[derive(Debug)]
18pub enum StreamInput<C: StreamCapability> {
19    /// One ordered Capability message.
20    Message(C::Message),
21    /// The consumer closed its sending direction while retaining its receive direction.
22    PeerHalfClosed,
23}
24
25enum ProviderOutput<C: StreamCapability> {
26    Message(C::Message),
27    PeerHalfClosed,
28    Terminal(Result<(), C::DomainError>),
29    Runtime(RuntimeFailure),
30}
31
32/// The typed provider side of one bounded bidirectional Stream session.
33///
34/// Plugin code uses this channel to exchange Capability values. Generated
35/// lowering keeps type erasure and [`NativeStreamSession`] behind the facade.
36pub 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    /// Sends one typed message with bounded backpressure.
58    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    /// Closes only the provider's sending direction.
66    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    /// Completes the Stream successfully.
75    pub async fn finish(&mut self) -> Result<(), RuntimeFailure> {
76        self.terminate(ProviderOutput::Terminal(Ok(()))).await
77    }
78
79    /// Completes the Stream with a Capability-defined Domain Error.
80    pub async fn fail(&mut self, error: C::DomainError) -> Result<(), RuntimeFailure> {
81        self.terminate(ProviderOutput::Terminal(Err(error))).await
82    }
83
84    /// Completes the Stream with an infrastructure Runtime Failure.
85    pub async fn fail_runtime(&mut self, error: RuntimeFailure) -> Result<(), RuntimeFailure> {
86        self.terminate(ProviderOutput::Runtime(error)).await
87    }
88
89    /// Closes the provider send direction and completes the Stream exactly once.
90    ///
91    /// Consuming the channel prevents Plugin code from accidentally sending or
92    /// terminating the session again after its operation result is known.
93    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    /// Receives the next typed consumer message or half-close marker.
108    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    /// Returns whether the consumer cancelled this Stream.
121    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
146/// One typed provider Stream erased to the native Adapter only after authoring.
147pub 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    /// Creates one bounded typed Stream tied to the invocation's cancellation.
158    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}