Skip to main content

dynamo_runtime/pipeline/nodes/
sources.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use super::*;
5use crate::pipeline::{AsyncEngine, PipelineIO};
6
7mod base;
8mod common;
9
10pub struct Frontend<In: PipelineIO, Out: PipelineIO> {
11    edge: OnceLock<Edge<In>>,
12    sinks: Arc<Mutex<HashMap<String, oneshot::Sender<Out>>>>,
13}
14
15/// A [`ServiceFrontend`] is the interface for an [`AsyncEngine<SingleIn<Context<In>>, ManyOut<Annotated<Out>>, Error>`]
16pub struct ServiceFrontend<In: PipelineIO, Out: PipelineIO> {
17    inner: Frontend<In, Out>,
18    self_weak: Weak<Self>,
19}
20
21pub struct SegmentSource<In: PipelineIO, Out: PipelineIO> {
22    inner: Frontend<In, Out>,
23    self_weak: Weak<Self>,
24}
25
26// impl<In: DataType, Out: PipelineIO> Frontend<In, Out> {
27//     pub fn new() -> Arc<Self> {
28//         Arc::new(Self {
29//             edge: OnceLock::new(),
30//             sinks: Arc::new(Mutex::new(HashMap::new())),
31//         })
32//     }
33// }
34
35// impl<In: DataType, Out: PipelineIO> SegmentSource<In, Out> {
36//     pub fn new() -> Arc<Self> {
37//         Arc::new(Self {
38//             edge: OnceLock::new(),
39//             sinks: Arc::new(Mutex::new(HashMap::new())),
40//         })
41//     }
42// }
43
44// #[async_trait]
45// impl<In: DataType, Out: PipelineIO> Source<Context<In>> for Frontend<In, Out> {
46//     async fn on_next(&self, data: Context<In>, _: private::Token) -> Result<(), PipelineError> {
47//         self.edge
48//             .get()
49//             .ok_or(PipelineError::NoEdge)?
50//             .write(data)
51//             .await
52//     }
53
54//     fn set_edge(
55//         &self,
56//         edge: Edge<Context<In>>>,
57//         _: private::Token,
58//     ) -> Result<(), PipelineError> {
59//         self.edge
60//             .set(edge)
61//             .map_err(|_| PipelineError::EdgeAlreadySet)?;
62//         Ok(())
63//     }
64// }
65
66// #[async_trait]
67// impl<In: DataType, Out: PipelineIO> Sink<PipelineStream<Out>> for Frontend<In, Out> {
68//     async fn on_data(
69//         &self,
70//         data: PipelineStream<Out>,
71//         _: private::Token,
72//     ) -> Result<(), PipelineError> {
73//         let context = data.context();
74
75//         let mut sinks = self.sinks.lock().unwrap();
76//         let tx = sinks
77//             .remove(context.id())
78//             .ok_or(PipelineError::DetachedStreamReceiver)
79//             .map_err(|e| {
80//                 data.context().stop_generating();
81//                 e
82//             })?;
83//         drop(sinks);
84
85//         let ctx = data.context();
86//         tx.send(data)
87//             .map_err(|_| PipelineError::DetachedStreamReceiver)
88//             .map_err(|e| {
89//                 ctx.stop_generating();
90//                 e
91//             })
92//     }
93// }
94
95// impl<In: DataType, Out: PipelineIO> Link<Context<In>> for Frontend<In, Out> {
96//     fn link<S: Sink<Context<In>> + 'static>(&self, sink: Arc<S>) -> Result<Arc<S>, PipelineError> {
97//         let edge = Edge::new(sink.clone());
98//         self.set_edge(edge.into(), private::Token {})?;
99//         Ok(sink)
100//     }
101// }
102
103// #[async_trait]
104// impl<In: DataType, Out: PipelineIO> AsyncEngine<Context<In>, Annotated<Out>, PipelineError>
105//     for Frontend<In, Out>
106// {
107//     async fn generate(&self, request: Context<In>) -> Result<PipelineStream<Out>, PipelineError> {
108//         let (tx, rx) = oneshot::channel::<PipelineStream<Out>>();
109//         {
110//             let mut sinks = self.sinks.lock().unwrap();
111//             sinks.insert(request.id().to_string(), tx);
112//         }
113//         self.on_next(request, private::Token {}).await?;
114//         rx.await.map_err(|_| PipelineError::DetachedStreamSender)
115//     }
116// }
117
118// // SegmentSource
119
120// #[async_trait]
121// impl<In: DataType, Out: PipelineIO> Source<Context<In>> for SegmentSource<In, Out> {
122//     async fn on_next(&self, data: Context<In>, _: private::Token) -> Result<(), PipelineError> {
123//         self.edge
124//             .get()
125//             .ok_or(PipelineError::NoEdge)?
126//             .write(data)
127//             .await
128//     }
129
130//     fn set_edge(
131//         &self,
132//         edge: Edge<Context<In>>>,
133//         _: private::Token,
134//     ) -> Result<(), PipelineError> {
135//         self.edge
136//             .set(edge)
137//             .map_err(|_| PipelineError::EdgeAlreadySet)?;
138//         Ok(())
139//     }
140// }
141
142// #[async_trait]
143// impl<In: DataType, Out: PipelineIO> Sink<PipelineStream<Out>> for SegmentSource<In, Out> {
144//     async fn on_data(
145//         &self,
146//         data: PipelineStream<Out>,
147//         _: private::Token,
148//     ) -> Result<(), PipelineError> {
149//         let context = data.context();
150
151//         let mut sinks = self.sinks.lock().unwrap();
152//         let tx = sinks
153//             .remove(context.id())
154//             .ok_or(PipelineError::DetachedStreamReceiver)
155//             .map_err(|e| {
156//                 data.context().stop_generating();
157//                 e
158//             })?;
159//         drop(sinks);
160
161//         let ctx = data.context();
162//         tx.send(data)
163//             .map_err(|_| PipelineError::DetachedStreamReceiver)
164//             .map_err(|e| {
165//                 ctx.stop_generating();
166//                 e
167//             })
168//     }
169// }
170
171// impl<In: DataType, Out: PipelineIO> Link<Context<In>> for SegmentSource<In, Out> {
172//     fn link<S: Sink<Context<In>> + 'static>(&self, sink: Arc<S>) -> Result<Arc<S>, PipelineError> {
173//         let edge = Edge::new(sink.clone());
174//         self.set_edge(edge.into(), private::Token {})?;
175//         Ok(sink)
176//     }
177// }
178
179// #[async_trait]
180// impl<In: DataType, Out: PipelineIO> AsyncEngine<Context<In>, Annotated<Out>, PipelineError>
181//     for SegmentSource<In, Out>
182// {
183//     async fn generate(&self, request: Context<In>) -> Result<PipelineStream<Out>, PipelineError> {
184//         let (tx, rx) = oneshot::channel::<PipelineStream<Out>>();
185//         {
186//             let mut sinks = self.sinks.lock().unwrap();
187//             sinks.insert(request.id().to_string(), tx);
188//         }
189//         self.on_next(request, private::Token {}).await?;
190//         rx.await.map_err(|_| PipelineError::DetachedStreamSender)
191//     }
192// }
193
194// #[cfg(test)]
195
196// mod tests {
197//     use super::*;
198
199//     #[tokio::test]
200//     async fn test_pipeline_source_no_edge() {
201//         let source = Frontend::<(), ()>::new();
202//         let stream = source.generate(().into()).await;
203//         match stream {
204//             Err(PipelineError::NoEdge) => (),
205//             _ => panic!("Expected NoEdge error"),
206//         }
207//     }
208// }
209
210// pub struct IngressPort<In, Out: PipelineIO> {
211//     edge: OnceLock<ServiceEngine<In, Out>>,
212// }
213
214// impl<In, Out> IngressPort<In, Out>
215// where
216//     In: for<'de> Deserialize<'de> + DataType,
217//     Out: PipelineIO + Serialize,
218// {
219//     pub fn new() -> Arc<Self> {
220//         Arc::new(IngressPort {
221//             edge: OnceLock::new(),
222//         })
223//     }
224// }
225
226// #[async_trait]
227// impl<In, Out> AsyncEngine<Context<Vec<u8>>, Vec<u8>> for IngressPort<In, Out>
228// where
229//     In: for<'de> Deserialize<'de> + DataType,
230//     Out: PipelineIO + Serialize,
231// {
232//     async fn generate(
233//         &self,
234//         request: Context<Vec<u8>>,
235//     ) -> Result<EngineStream<Vec<u8>>, PipelineError> {
236//         // Deserialize request
237//         let request = request.try_map(|bytes| {
238//             bincode::deserialize::<In>(&bytes)
239//                 .map_err(|err| PipelineError(format!("Failed to deserialize request: {}", err)))
240//         })?;
241
242//         // Forward request to edge
243//         let stream = self
244//             .edge
245//             .get()
246//             .ok_or(PipelineError("No engine to forward request to".to_string()))?
247//             .generate(request)
248//             .await?;
249
250//         // Serialize response stream
251
252//         let stream =
253//             stream.map(|resp| bincode::serialize(&resp).expect("Failed to serialize response"));
254
255//         Err(PipelineError(format!("Not implemented")))
256//     }
257// }
258
259// fn convert_stream<T, U>(
260//     stream: impl Stream<Item = ServerStream<T>> + Send + 'static,
261//     ctx: Arc<dyn AsyncEngineContext>,
262//     transform: Arc<dyn Fn(T) -> Result<U, StreamError> + Send + Sync>,
263// ) -> Pin<Box<dyn Stream<Item = ServerStream<U>> + Send>>
264// where
265//     T: Send + 'static,
266//     U: Send + 'static,
267// {
268//     Box::pin(stream.flat_map(move |item| {
269//         let ctx = ctx.clone();
270//         let transform = transform.clone();
271//         match item {
272//             ServerStream::Data(data) => match transform(data) {
273//                 Ok(transformed) => futures::stream::iter(vec![ServerStream::Data(transformed)]),
274//                 Err(e) => {
275//                     // Trigger cancellation and propagate the error, followed by Sentinel
276//                     ctx.stop_generating();
277//                     futures::stream::iter(vec![ServerStream::Error(e), ServerStream::Sentinel])
278//                 }
279//             },
280//             other => futures::stream::iter(vec![other]),
281//         }
282//     })
283//     // Use take_while to stop processing when encountering the Sentinel
284//     .take_while(|item| futures::future::ready(!matches!(item, ServerStream::Sentinel))))
285// }