Skip to main content

rskit_dataset/
stream.rs

1//! Stream adapters for composing dataset items with `rskit-stream`.
2
3use futures::Stream;
4use rskit_errors::AppResult;
5use rskit_stream::RskitStreamExt;
6
7use crate::{DataItem, DatasetLimits, Transform};
8
9/// Extension methods for streams of dataset items.
10pub trait DatasetStreamExt: Stream<Item = AppResult<DataItem>> + Sized + Send + 'static {
11    /// Apply a fallible dataset transform inside a canonical `rskit-stream` stream.
12    fn apply_dataset_transform<T>(
13        self,
14        transform: T,
15        limits: DatasetLimits,
16    ) -> impl Stream<Item = AppResult<Option<DataItem>>> + Send + 'static
17    where
18        T: Transform<DataItem, DataItem> + Clone + Send + Sync + 'static,
19    {
20        self.rmap(move |item| {
21            let transform = transform.clone();
22            async move {
23                let item = item?;
24                transform.apply(item, &limits)
25            }
26        })
27    }
28}
29
30impl<S> DatasetStreamExt for S where S: Stream<Item = AppResult<DataItem>> + Sized + Send + 'static {}
31
32#[cfg(test)]
33mod tests {
34    use futures::{StreamExt, stream};
35    use rskit_errors::{AppError, ErrorCode};
36
37    use super::*;
38    use crate::{DataItem, Label, MediaType};
39
40    #[derive(Clone)]
41    struct FilterBySource {
42        allowed: &'static str,
43    }
44
45    impl Transform<DataItem, DataItem> for FilterBySource {
46        fn name(&self) -> &str {
47            "filter-by-source"
48        }
49
50        fn apply(&self, item: DataItem, _limits: &DatasetLimits) -> AppResult<Option<DataItem>> {
51            if item.source_name == self.allowed {
52                Ok(Some(item))
53            } else {
54                Ok(None)
55            }
56        }
57    }
58
59    #[tokio::test]
60    async fn dataset_stream_transform_maps_items_and_forwards_errors() {
61        let transform = FilterBySource { allowed: "keep" };
62        assert_eq!(transform.name(), "filter-by-source");
63        let keep = DataItem::new(vec![1], Label::Real, MediaType::Text, "keep").unwrap();
64        let drop = DataItem::new(vec![2], Label::AiGenerated, MediaType::Text, "drop").unwrap();
65        let input = stream::iter([
66            Ok(keep),
67            Ok(drop),
68            Err(AppError::new(ErrorCode::Internal, "boom")),
69        ]);
70
71        let output = input
72            .apply_dataset_transform(transform, DatasetLimits::default())
73            .collect::<Vec<_>>()
74            .await;
75
76        assert!(output[0].as_ref().unwrap().is_some());
77        assert!(output[1].as_ref().unwrap().is_none());
78        assert_eq!(output[2].as_ref().unwrap_err().code(), ErrorCode::Internal);
79    }
80}