Skip to main content

lumen_engine/media/
prediction.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use crate::{
4    composition::Composition,
5    error::{GraphValidationError, LumenError, MediaError},
6    expr::ExpressionContext,
7    node::{
8        NodeId, NodeKind, NodeParamEvalContext, NodeParams, PortRef,
9        processing::time_remap::{TimeRemapSettings, remap_frame},
10        source::media_in,
11    },
12};
13
14use super::MediaStore;
15
16#[derive(Debug, Clone, PartialEq, Eq, Default)]
17pub struct FrameRequirements {
18    pub images: Vec<String>,
19    pub videos: Vec<VideoFrameRequirement>,
20}
21
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct VideoFrameRequirement {
24    pub stream_id: String,
25    pub frames: Vec<u32>,
26}
27
28#[derive(Debug, Clone, Default)]
29pub struct RenderRequirements {
30    images: BTreeSet<String>,
31    videos: BTreeMap<String, BTreeSet<u32>>,
32}
33
34impl RenderRequirements {
35    pub fn add_image(&mut self, image_id: impl Into<String>) {
36        self.images.insert(image_id.into());
37    }
38
39    pub fn add_video_frame(&mut self, stream_id: impl Into<String>, frame: u32) {
40        self.videos
41            .entry(stream_id.into())
42            .or_default()
43            .insert(frame);
44    }
45
46    pub fn merge(&mut self, other: FrameRequirements) {
47        self.images.extend(other.images);
48        for video in other.videos {
49            self.videos
50                .entry(video.stream_id)
51                .or_default()
52                .extend(video.frames);
53        }
54    }
55}
56
57impl From<RenderRequirements> for FrameRequirements {
58    fn from(value: RenderRequirements) -> Self {
59        Self {
60            images: value.images.into_iter().collect(),
61            videos: value
62                .videos
63                .into_iter()
64                .map(|(stream_id, frames)| VideoFrameRequirement {
65                    stream_id,
66                    frames: frames.into_iter().collect(),
67                })
68                .collect(),
69        }
70    }
71}
72
73pub fn collect_frame_requirements<M: MediaStore>(
74    composition: &Composition,
75    media_store: &M,
76    frame: u32,
77) -> Result<FrameRequirements, LumenError> {
78    tracing::trace!(target: "lumen_media", frame, "collect frame requirements");
79    let output_port = media_output_port(composition)?;
80    let mut collector = RenderRequirements::default();
81    let mut context = RequirementContext {
82        composition,
83        media_store,
84        frame,
85    };
86    context.collect_port(&output_port, &mut collector)?;
87    let requirements = FrameRequirements::from(collector);
88    tracing::trace!(
89        target: "lumen_media",
90        frame,
91        images = requirements.images.len(),
92        videos = requirements.videos.len(),
93        "collected frame requirements"
94    );
95    Ok(requirements)
96}
97
98struct RequirementContext<'a, M: MediaStore> {
99    composition: &'a Composition,
100    media_store: &'a M,
101    frame: u32,
102}
103
104impl<'a, M: MediaStore> RequirementContext<'a, M> {
105    fn collect_port(
106        &mut self,
107        port: &PortRef,
108        collector: &mut RenderRequirements,
109    ) -> Result<(), LumenError> {
110        if port.is_empty() {
111            return Ok(());
112        }
113
114        let Some(node) = self.composition.graph.nodes.get(&port.id) else {
115            return Ok(());
116        };
117        self.collect_node(port.id, node, collector)
118    }
119
120    fn collect_node(
121        &mut self,
122        node_id: NodeId,
123        node: &NodeKind,
124        collector: &mut RenderRequirements,
125    ) -> Result<(), LumenError> {
126        match node {
127            NodeKind::MediaOutput(media_output) => {
128                self.collect_port(&media_output.source, collector)?;
129            }
130            NodeKind::MediaIn(media_in_node) => {
131                self.collect_media_in(media_in_node, collector)?;
132            }
133            NodeKind::TimeRemap(time_remap) => {
134                let target_frame = self.remapped_frame(time_remap)?;
135                self.with_frame(target_frame, |context| {
136                    context.collect_port(&time_remap.source, collector)
137                })?;
138            }
139            NodeKind::Switch(switch) => {
140                if let Some(layer) = crate::node::compositing::switch::selected_layer_for_frame(
141                    switch,
142                    &self.expr_context("switch_requirements"),
143                )?
144                .and_then(|index| switch.layers.get(index))
145                {
146                    self.collect_port(layer, collector)?;
147                }
148            }
149            _ => self.collect_default_inputs(node_id, collector)?,
150        }
151
152        Ok(())
153    }
154
155    fn collect_default_inputs(
156        &mut self,
157        node_id: NodeId,
158        collector: &mut RenderRequirements,
159    ) -> Result<(), LumenError> {
160        let inputs: Vec<_> = self
161            .composition
162            .graph
163            .connections
164            .iter()
165            .filter(|connection| connection.to_node == node_id)
166            .map(|connection| PortRef::new(connection.from_node, connection.from_port.clone()))
167            .collect();
168
169        for input in inputs {
170            self.collect_port(&input, collector)?;
171        }
172
173        Ok(())
174    }
175
176    fn collect_media_in(
177        &self,
178        media_in_node: &media_in::MediaIn,
179        collector: &mut RenderRequirements,
180    ) -> Result<(), LumenError> {
181        match media_in::resolve_for_context(
182            media_in_node,
183            &self.expr_context("media_requirements"),
184        )? {
185            media_in::MediaInKind::Image { image_id } => {
186                tracing::trace!(
187                    target: "lumen_media",
188                    frame = self.frame,
189                    image_id = %image_id,
190                    "require image"
191                );
192                collector.add_image(image_id);
193            }
194            media_in::MediaInKind::Video {
195                stream_id,
196                range,
197                speed,
198                loop_mode,
199            } => {
200                let resolver =
201                    self.media_store
202                        .get_video_resolver(&stream_id)
203                        .ok_or_else(|| MediaError::SourceNotFound {
204                            media_source: stream_id.clone(),
205                        })?;
206                let metadata = resolver.metadata();
207                let source_frame = media_in::map_to_source_frame(
208                    self.frame,
209                    self.composition.timeline.fps,
210                    metadata.fps,
211                    metadata.frame_count,
212                    range.as_ref(),
213                    speed,
214                    loop_mode,
215                )
216                .ok_or_else(|| MediaError::FrameOutOfRange {
217                    media_source: stream_id.clone(),
218                    frame: self.frame,
219                    frame_count: metadata.frame_count,
220                })?;
221                tracing::trace!(
222                    target: "lumen_media",
223                    frame = self.frame,
224                    stream_id = %stream_id,
225                    source_frame,
226                    "require video frame"
227                );
228                collector.add_video_frame(stream_id, source_frame);
229            }
230        }
231
232        Ok(())
233    }
234
235    fn remapped_frame(
236        &self,
237        time_remap: &crate::node::processing::time_remap::TimeRemap,
238    ) -> Result<u32, LumenError> {
239        let expr_context = self.expr_context("time_remap_requirements");
240        let params = time_remap.params.eval(&NodeParamEvalContext {
241            node_id: time_remap.id,
242            expr: &expr_context,
243        })?;
244        Ok(remap_frame(TimeRemapSettings {
245            frame: params.frame,
246            loop_enabled: params.loop_enabled,
247            loop_start: params.loop_start,
248            loop_end: params.loop_end,
249        }))
250    }
251
252    fn with_frame<T>(
253        &mut self,
254        frame: u32,
255        f: impl FnOnce(&mut Self) -> Result<T, LumenError>,
256    ) -> Result<T, LumenError> {
257        let original_frame = self.frame;
258        self.frame = frame;
259        let result = f(self);
260        self.frame = original_frame;
261        result
262    }
263
264    fn expr_context(&self, path: &str) -> ExpressionContext<'_> {
265        ExpressionContext {
266            frame: self.frame,
267            fps: self.composition.timeline.fps,
268            width: self.composition.render_settings.width,
269            height: self.composition.render_settings.height,
270            duration_frames: self.composition.timeline.duration_frames,
271            path: Some(path.to_string()),
272            graph: Some(&self.composition.graph),
273        }
274    }
275}
276
277fn media_output_port(composition: &Composition) -> Result<PortRef, LumenError> {
278    let mut media_outputs = composition
279        .graph
280        .nodes
281        .iter()
282        .filter_map(|(node_id, node)| matches!(node, NodeKind::MediaOutput(_)).then_some(*node_id));
283    let Some(output_node_id) = media_outputs.next() else {
284        return Err(GraphValidationError::MissingMediaOutput.into());
285    };
286    if media_outputs.next().is_some() {
287        return Err(GraphValidationError::MultipleMediaOutputs { count: 2 }.into());
288    }
289
290    Ok(PortRef::new(output_node_id, "output".to_string()))
291}