1use futures::Stream;
4use rskit_errors::AppResult;
5use rskit_stream::RskitStreamExt;
6
7use crate::{DataItem, DatasetLimits, Transform};
8
9pub trait DatasetStreamExt: Stream<Item = AppResult<DataItem>> + Sized + Send + 'static {
11 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}