Skip to main content

gstreamer_app/
app_src_futures.rs

1// Take a look at the license at the top of the repository in the LICENSE file.
2
3use crate::{AppSrc, AppSrcCallbacks};
4use futures_sink::Sink;
5use glib::object::ObjectExt as _;
6use std::{
7    pin::Pin,
8    sync::{Arc, Mutex},
9    task::Waker,
10    task::{Context, Poll},
11};
12
13#[derive(Debug)]
14pub struct AppSrcSink {
15    app_src: glib::WeakRef<AppSrc>,
16    waker_reference: Arc<Mutex<Option<Waker>>>,
17}
18
19impl AppSrcSink {
20    pub(crate) fn new(app_src: &AppSrc) -> Self {
21        skip_assert_initialized!();
22
23        let waker_reference = Arc::new(Mutex::new(None as Option<Waker>));
24
25        app_src.set_callbacks(
26            AppSrcCallbacks::builder()
27                .need_data({
28                    let waker_reference = Arc::clone(&waker_reference);
29
30                    move |_, _| {
31                        if let Some(waker) = waker_reference.lock().unwrap().take() {
32                            waker.wake();
33                        }
34                    }
35                })
36                .build(),
37        );
38
39        Self {
40            app_src: app_src.downgrade(),
41            waker_reference,
42        }
43    }
44}
45
46impl Drop for AppSrcSink {
47    fn drop(&mut self) {
48        #[cfg(not(feature = "v1_18"))]
49        {
50            // This is not thread-safe before 1.16.3, see
51            // https://gitlab.freedesktop.org/gstreamer/gst-plugins-base/merge_requests/570
52            if gst::version() >= (1, 16, 3, 0)
53                && let Some(app_src) = self.app_src.upgrade()
54            {
55                app_src.set_callbacks(AppSrcCallbacks::builder().build());
56            }
57        }
58    }
59}
60
61impl Sink<gst::Sample> for AppSrcSink {
62    type Error = gst::FlowError;
63
64    fn poll_ready(self: Pin<&mut Self>, context: &mut Context) -> Poll<Result<(), Self::Error>> {
65        let mut waker = self.waker_reference.lock().unwrap();
66
67        let Some(app_src) = self.app_src.upgrade() else {
68            return Poll::Ready(Err(gst::FlowError::Eos));
69        };
70
71        let current_level_bytes = app_src.current_level_bytes();
72        let max_bytes = app_src.max_bytes();
73
74        if current_level_bytes >= max_bytes && max_bytes != 0 {
75            waker.replace(context.waker().to_owned());
76
77            Poll::Pending
78        } else {
79            Poll::Ready(Ok(()))
80        }
81    }
82
83    fn start_send(self: Pin<&mut Self>, sample: gst::Sample) -> Result<(), Self::Error> {
84        let Some(app_src) = self.app_src.upgrade() else {
85            return Err(gst::FlowError::Eos);
86        };
87
88        app_src.push_sample(&sample)?;
89
90        Ok(())
91    }
92
93    fn poll_flush(self: Pin<&mut Self>, _: &mut Context) -> Poll<Result<(), Self::Error>> {
94        Poll::Ready(Ok(()))
95    }
96
97    fn poll_close(self: Pin<&mut Self>, _: &mut Context) -> Poll<Result<(), Self::Error>> {
98        let Some(app_src) = self.app_src.upgrade() else {
99            return Poll::Ready(Ok(()));
100        };
101
102        app_src.end_of_stream()?;
103
104        Poll::Ready(Ok(()))
105    }
106}
107
108#[cfg(test)]
109mod tests {
110    use futures_util::{sink::SinkExt, stream::StreamExt};
111    use gst::prelude::*;
112    use std::sync::{
113        Arc,
114        atomic::{AtomicUsize, Ordering},
115    };
116
117    use super::*;
118
119    #[test]
120    fn test_app_src_sink() {
121        gst::init().unwrap();
122
123        let appsrc = gst::ElementFactory::make("appsrc").build().unwrap();
124        let fakesink = gst::ElementFactory::make("fakesink")
125            .property("signal-handoffs", true)
126            .build()
127            .unwrap();
128
129        let pipeline = gst::Pipeline::new();
130        pipeline.add(&appsrc).unwrap();
131        pipeline.add(&fakesink).unwrap();
132
133        appsrc.link(&fakesink).unwrap();
134
135        let mut bus_stream = pipeline.bus().unwrap().stream();
136        let mut app_src_sink = appsrc.dynamic_cast::<AppSrc>().unwrap().sink();
137
138        let sample_quantity = 5;
139
140        let samples = (0..sample_quantity)
141            .map(|_| gst::Sample::builder().buffer(&gst::Buffer::new()).build())
142            .collect::<Vec<gst::Sample>>();
143
144        let mut sample_stream = futures_util::stream::iter(samples).map(Ok);
145
146        let handoff_count_reference = Arc::new(AtomicUsize::new(0));
147
148        fakesink.connect("handoff", false, {
149            let handoff_count_reference = Arc::clone(&handoff_count_reference);
150
151            move |_| {
152                handoff_count_reference.fetch_add(1, Ordering::AcqRel);
153
154                None
155            }
156        });
157
158        pipeline.set_state(gst::State::Playing).unwrap();
159
160        futures_executor::block_on(app_src_sink.send_all(&mut sample_stream)).unwrap();
161        futures_executor::block_on(app_src_sink.close()).unwrap();
162
163        while let Some(message) = futures_executor::block_on(bus_stream.next()) {
164            match message.view() {
165                gst::MessageView::Eos(_) => break,
166                gst::MessageView::Error(_) => unreachable!(),
167                _ => continue,
168            }
169        }
170
171        pipeline.set_state(gst::State::Null).unwrap();
172
173        assert_eq!(
174            handoff_count_reference.load(Ordering::Acquire),
175            sample_quantity
176        );
177    }
178}