Skip to main content

dynamo_runtime/pipeline/nodes/sources/
common.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use crate::engine::AsyncEngineContextProvider;
5
6use super::*;
7
8macro_rules! impl_frontend {
9    ($type:ident) => {
10        impl<In: PipelineIO, Out: PipelineIO> $type<In, Out> {
11            pub fn new() -> Arc<Self> {
12                Arc::new_cyclic(|self_weak| Self {
13                    inner: Frontend::default(),
14                    self_weak: self_weak.clone(),
15                })
16            }
17        }
18
19        #[async_trait]
20        impl<In: PipelineIO, Out: PipelineIO> Source<In> for $type<In, Out> {
21            async fn on_next(&self, data: In, token: private::Token) -> Result<(), Error> {
22                self.inner.on_next(data, token).await
23            }
24
25            fn set_edge(&self, edge: Edge<In>, token: private::Token) -> Result<(), PipelineError> {
26                self.inner.set_edge(edge, token)
27            }
28        }
29
30        #[async_trait]
31        impl<In: PipelineIO, Out: PipelineIO + AsyncEngineContextProvider> Sink<Out>
32            for $type<In, Out>
33        {
34            async fn on_data(&self, data: Out, token: private::Token) -> Result<(), Error> {
35                self.inner.on_data(data, token).await
36            }
37        }
38
39        #[async_trait]
40        impl<In: PipelineIO + Sync, Out: PipelineIO> AsyncEngine<In, Out, Error>
41            for $type<In, Out>
42        {
43            async fn generate(&self, request: In) -> Result<Out, Error> {
44                if let Some(root) = self.self_weak.upgrade() {
45                    request.context().retain(root.clone());
46                    let response = self.inner.generate(request).await?;
47                    response.context().retain(root);
48                    Ok(response)
49                } else {
50                    self.inner.generate(request).await
51                }
52            }
53        }
54    };
55}
56
57impl_frontend!(ServiceFrontend);
58impl_frontend!(SegmentSource);
59
60#[cfg(test)]
61mod tests {
62    use super::*;
63    use crate::pipeline::{ManyOut, PipelineErrorExt, SingleIn};
64
65    #[tokio::test]
66    async fn test_pipeline_source_no_edge() {
67        let source = Frontend::<SingleIn<()>, ManyOut<()>>::default();
68        let stream = source
69            .generate(().into())
70            .await
71            .unwrap_err()
72            .try_into_pipeline_error()
73            .unwrap();
74
75        match stream {
76            PipelineError::NoEdge => (),
77            _ => panic!("Expected NoEdge error"),
78        }
79    }
80}