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// }