Skip to main content

utsuru/mirrors/discord/
mod.rs

1use bytes::Bytes;
2use davey::{Codec, DaveSession, MediaType};
3use serde_json::json;
4use std::{
5    collections::HashSet,
6    error::Error as StdError,
7    fmt::{Display, Formatter, Result as FmtResult},
8    num::ParseIntError,
9    pin::Pin,
10    sync::{
11        Arc,
12        atomic::{AtomicBool, Ordering},
13    },
14    time::Duration,
15};
16use tokio::{
17    sync::{
18        Notify, RwLock,
19        mpsc::{self, error::SendError},
20        oneshot::{self, error::RecvError},
21    },
22    time::sleep,
23};
24use tokio_websockets::{Message as WebSocketMessage, Payload};
25use tracing::debug;
26use twilight_gateway::{Intents, Shard, ShardId};
27use twilight_model::id::{
28    Id,
29    marker::{ChannelMarker, GuildMarker},
30};
31use uuid::Uuid;
32use webrtc::{
33    api::media_engine::{MIME_TYPE_H264, MIME_TYPE_OPUS},
34    media::Sample,
35    peer_connection::sdp::{sdp_type::RTCSdpType, session_description::RTCSessionDescription},
36    rtp_transceiver::rtp_codec::RTCRtpCodecCapability,
37    track::track_local::{TrackLocal, track_local_static_sample::TrackLocalStaticSample},
38};
39
40use super::Mirror;
41use crate::error::{Error, ErrorType};
42use crate::utils::{h264_parser::parse_sps, h264_synthesizer::synthesize_sps};
43
44mod dave;
45mod endpoint;
46mod gateway;
47mod heartbeat;
48
49const NALU_SHORT_START_SEQUENCE_SIZE: usize = 3;
50const START_CODE_HIGHEST_POSSIBLE_VALUE: u8 = 1;
51const START_CODE_END_BYTE_VALUE: u8 = 1;
52const START_CODE_LEADING_BYTES_VALUE: u8 = 0;
53
54pub struct DiscordLiveBuilder {
55    token: Box<str>,
56    guild_id: Id<GuildMarker>,
57    channel_id: Id<ChannelMarker>,
58}
59
60impl DiscordLiveBuilder {
61    pub fn new(token: impl AsRef<str>, guild_id: u64, channel_id: u64) -> Self {
62        Self {
63            token: token.as_ref().into(),
64            guild_id: Id::new(guild_id),
65            channel_id: Id::new(channel_id),
66        }
67    }
68
69    pub async fn connect(
70        self,
71        trace_tx: Option<mpsc::UnboundedSender<DiscordLiveBuilderState>>,
72    ) -> Result<DiscordLive, Error<dyn ErrorInner>> {
73        let _ = rustls::crypto::ring::default_provider().install_default();
74
75        if self.token.len() < 4 {
76            return Err(Error {
77                kind: ErrorType::DiscordAuth,
78                source: None,
79            });
80        }
81
82        let mut token = String::from(self.token.as_ref());
83        token.replace_range(0..4, "Bot ");
84        let token_ptr: *mut u8 = token.as_mut_ptr();
85
86        let intents =
87            Intents::GUILD_MESSAGES | Intents::GUILD_VOICE_STATES | Intents::MESSAGE_CONTENT;
88        let shard = Shard::new(ShardId::ONE, token, intents);
89
90        let src = self.token.as_bytes();
91        unsafe {
92            *token_ptr = src[0];
93            *token_ptr.add(1) = src[1];
94            *token_ptr.add(2) = src[2];
95            *token_ptr.add(3) = src[3];
96        }
97
98        let (voice_tx, voice_rx) = oneshot::channel();
99        let voice_tx = Some(voice_tx);
100        let (rtcsrv_tx, rtcsrv_rx) = oneshot::channel();
101        let rtcsrv_tx = Some(rtcsrv_tx);
102        let (wsconn_tx, wsconn_rx) = oneshot::channel();
103        let wsconn_tx = Some(wsconn_tx);
104        let (feed_tx, feed_rx) = oneshot::channel();
105        let (nego_tx, nego_rx) = oneshot::channel();
106        let nego_tx = Some(nego_tx);
107        let (connected_tx, connected_rx) = oneshot::channel();
108        let connected_tx = Some(connected_tx);
109        let (remote_tx, remote_rx) = oneshot::channel();
110        let remote_tx = Some(remote_tx);
111        let (heartbeat_tx, heartbeat_rx) = oneshot::channel();
112        let heartbeat_tx = Some(heartbeat_tx);
113        let (instance_tx, instance_rx) = oneshot::channel();
114        let instance_tx = Some(instance_tx);
115        let (egress_tx, egress_rx) = mpsc::unbounded_channel();
116        let (nonce_tx, nonce_rx) = mpsc::unbounded_channel();
117        let (dave_tx, dave_rx) = mpsc::unbounded_channel();
118
119        let audio_payload = 111;
120        let audio_codec = "opus";
121        let mut audio_mid: u8 = 0;
122        let mut audio_ssrc: u32 = 0;
123        let video_payload = 102;
124        let video_codec = "H264";
125        let video_rtxpayload = 103;
126        let mut video_mid: u8 = 1;
127        let mut video_ssrc: u32 = 0;
128        let mut video_rtxssrc: u32 = 0;
129
130        let notify = Arc::new(Notifier::new());
131
132        if let Err(e) = gateway::handle(&notify, self, shard, voice_tx, rtcsrv_tx, wsconn_tx).await
133        {
134            notify.close();
135            return Err(Error {
136                kind: e.kind,
137                source: e.source.map(|source| source as Box<dyn ErrorInner>),
138            });
139        }
140
141        trace_tx
142            .as_ref()
143            .map(|tx| tx.send(DiscordLiveBuilderState::VoiceConnecting));
144        let (user_id, session_id) = voice_rx.await?;
145        trace_tx
146            .as_ref()
147            .map(|tx| tx.send(DiscordLiveBuilderState::StreamCreating));
148        let (server, channel) = rtcsrv_rx.await?;
149        let channel_id: Result<u64, _> = channel.parse();
150        trace_tx
151            .as_ref()
152            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointWSConnecting));
153        let (token, endpoint) = wsconn_rx.await?;
154        if let Err(e) = endpoint::handle(
155            &notify,
156            user_id.to_string(),
157            session_id,
158            server,
159            channel,
160            token,
161            endpoint,
162            audio_payload,
163            audio_codec,
164            video_payload,
165            video_codec,
166            video_rtxpayload,
167            egress_rx,
168            feed_tx,
169            nego_tx,
170            connected_tx,
171            remote_tx,
172            nonce_tx,
173            heartbeat_tx,
174            &dave_tx,
175        )
176        .await
177        {
178            notify.close();
179            return Err(Error {
180                kind: e.kind,
181                source: e.source.map(|source| source as Box<dyn ErrorInner>),
182            });
183        }
184
185        trace_tx
186            .as_ref()
187            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointRTCCreating));
188        let (peer_connection, audio_rtp_sender, video_rtp_sender, streams) = feed_rx.await?;
189
190        let heartbeat_interval = heartbeat_rx.await?;
191        if let Err(e) = heartbeat::handle(&notify, heartbeat_interval, &egress_tx, nonce_rx).await {
192            notify.close();
193            return Err(Error {
194                kind: e.kind,
195                source: e.source.map(|source| source as Box<dyn ErrorInner>),
196            });
197        }
198
199        trace_tx
200            .as_ref()
201            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointRTCNegotiation));
202        nego_rx.await?;
203        if let Err(e) = dave::handle(&notify, &egress_tx, dave_rx, instance_tx).await {
204            notify.close();
205            return Err(Error {
206                kind: e.kind,
207                source: e.source.map(|source| source as Box<dyn ErrorInner>),
208            });
209        }
210
211        let offer = peer_connection.create_offer(None).await?;
212        let mut gather_complete = peer_connection.gathering_complete_promise().await;
213        peer_connection.set_local_description(offer).await?;
214        let _ = gather_complete.recv().await;
215        let local_desc = peer_connection.local_description().await.ok_or(Error {
216            kind: ErrorType::DiscordEndpoint,
217            source: None,
218        })?;
219
220        let sdp = local_desc.unmarshal()?;
221        let mut attributes = HashSet::new();
222        for attribute in sdp.attributes {
223            if attribute.key.as_str() == "fingerprint" {
224                if let Some(value) = attribute.value {
225                    attributes.insert(format!("a={}:{}", attribute.key, value));
226                } else {
227                    attributes.insert(format!("a={}", attribute.key));
228                }
229            }
230        }
231        for media in sdp.media_descriptions {
232            for attribute in media.attributes {
233                match attribute.key.as_str() {
234                    "ice-ufrag" | "ice-pwd" | "ice-options" | "extmap" | "rtpmap" => {
235                        if let Some(value) = attribute.value {
236                            attributes.insert(format!("a={}:{}", attribute.key, value));
237                        } else {
238                            attributes.insert(format!("a={}", attribute.key));
239                        }
240                    }
241                    "ssrc" => {
242                        if media.media_name.media.as_str() == "audio"
243                            && let Some(value) = attribute.value
244                        {
245                            let mut value = value.split_whitespace();
246                            audio_ssrc = value
247                                .next()
248                                .ok_or(Error {
249                                    kind: ErrorType::DiscordEndpoint,
250                                    source: None,
251                                })?
252                                .parse()?;
253                        }
254                    }
255                    "ssrc-group" => {
256                        if media.media_name.media.as_str() == "video"
257                            && let Some(value) = attribute.value
258                        {
259                            let mut value = value.split_whitespace();
260                            let _ = value.next();
261                            video_ssrc = value
262                                .next()
263                                .ok_or(Error {
264                                    kind: ErrorType::DiscordEndpoint,
265                                    source: None,
266                                })?
267                                .parse()?;
268                            video_rtxssrc = value
269                                .next()
270                                .ok_or(Error {
271                                    kind: ErrorType::DiscordEndpoint,
272                                    source: None,
273                                })?
274                                .parse()?;
275                        }
276                    }
277                    "mid" => match media.media_name.media.as_str() {
278                        "audio" => {
279                            if let Some(value) = attribute.value {
280                                audio_mid = value
281                                    .split_whitespace()
282                                    .next()
283                                    .ok_or(Error {
284                                        kind: ErrorType::DiscordEndpoint,
285                                        source: None,
286                                    })?
287                                    .parse()?;
288                            }
289                        }
290                        "video" => {
291                            if let Some(value) = attribute.value {
292                                video_mid = value
293                                    .split_whitespace()
294                                    .next()
295                                    .ok_or(Error {
296                                        kind: ErrorType::DiscordEndpoint,
297                                        source: None,
298                                    })?
299                                    .parse()?;
300                            }
301                        }
302                        _ => {}
303                    },
304                    _ => {}
305                }
306            }
307        }
308        let attributes = attributes.into_iter().collect::<Vec<_>>().join("\n");
309
310        let sdp = format!("a=extmap-allow-mixed\n{}", attributes);
311        let payload = json!({
312            "op": 1,
313            "d": {
314                "protocol": "webrtc",
315                "data": sdp,
316                "sdp": sdp,
317                "codecs": [
318                    {"name": audio_codec, "type": "audio", "priority": 1000, "payload_type": audio_payload, "rtx_payload_type": null},
319                    {"name": video_codec, "type": "video", "priority": 1000, "payload_type": video_payload, "rtx_payload_type": video_rtxpayload}
320                ],
321                "rtc_connection_id": Uuid::new_v4().to_string()
322            }
323        });
324        egress_tx.send(WebSocketMessage::text(payload.to_string()))?;
325        debug!("[WebRTC] offer sent, waiting for answer");
326
327        trace_tx
328            .as_ref()
329            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointWSSDP));
330        let (remote_sdp, dave_protocol_version, external_payload) = remote_rx.await?;
331
332        let mut answer = RTCSessionDescription::default();
333        answer.sdp_type = RTCSdpType::Answer;
334        let remote_sdp = remote_sdp
335            .replace("ICE/SDP", &format!("UDP/TLS/RTP/SAVPF {audio_payload}"))
336            .replace("\n", "\r\n");
337        let remote_sdp = format!(
338            "v=0\r\no=- 1420070400000 0 IN IP4 127.0.0.1\r\ns=-\r\nt=0 0\r\na=msid-semantic: WMS *\r\na=group:BUNDLE 0 1\r\n\
339            {remote_sdp}"
340        );
341        answer.sdp = remote_sdp;
342
343        let parsed = answer.unmarshal()?;
344        let port = &parsed.media_descriptions[0].media_name.port.value;
345        let connection = &parsed.media_descriptions[0].connection_information;
346        let attributes = &parsed.media_descriptions[0].attributes;
347        let setup = "passive";
348        let direction = "inactive";
349        let remote_sdp = format!(
350            "v=0\r\no=- 1420070400000 0 IN IP4 127.0.0.1\r\ns=-\r\nt=0 0\r\na=msid-semantic: WMS *\r\na=group:BUNDLE 0 1\r\n\
351            m=audio {port} UDP/TLS/RTP/SAVPF {audio_payload}\r\na=rtpmap:{audio_payload} {audio_codec}/48000/2\r\na=fmtp:{audio_payload} minptime=10;useinbandfec=1;usedtx=0\r\na=rtcp-fb:{audio_payload} transport-cc\r\na=extmap:1 urn:ietf:params:rtp-hdrext:ssrc-audio-level\r\na=extmap:3 http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01\r\na=setup:{setup}\r\na=mid:{audio_mid}\r\na=maxptime:60\r\na={direction}\r\na=rtcp-mux\r\n\
352            m=video {port} UDP/TLS/RTP/SAVPF {video_payload} {video_rtxpayload}\r\na=rtpmap:{video_payload} {video_codec}/90000\r\na=rtpmap:{video_rtxpayload} rtx/90000\r\na=fmtp:{video_payload} x-google-max-bitrate=2500;level-asymmetry-allowed=1;packetization-mode=1;profile-level-id=42e01f\r\na=fmtp:{video_rtxpayload} apt={video_payload}\r\na=rtcp-fb:{video_payload} ccm fir\r\na=rtcp-fb:{video_payload} nack\r\na=rtcp-fb:{video_payload} nack pli\r\na=rtcp-fb:{video_payload} goog-remb\r\na=rtcp-fb:{video_payload} transport-cc\r\na=extmap:2 http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time\r\na=extmap:3 http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01\r\na=extmap:14 urn:ietf:params:rtp-hdrext:toffset\r\na=extmap:13 urn:3gpp:video-orientation\r\na=extmap:5 http://www.webrtc.org/experiments/rtp-hdrext/playout-delay\r\na=setup:{setup}\r\na=mid:{video_mid}\r\na={direction}\r\na=rtcp-mux\r\n"
353        );
354        answer.sdp = remote_sdp;
355
356        let mut parsed = answer.unmarshal()?;
357        for media in &mut parsed.media_descriptions {
358            media.connection_information = connection.clone();
359            for attribute in attributes {
360                media.attributes.push(attribute.clone());
361            }
362        }
363        let remote_sdp = parsed.marshal();
364        let inactive_sdp = RTCSessionDescription::answer(remote_sdp)?;
365
366        let direction = "recvonly";
367        let remote_sdp = format!(
368            "v=0\r\no=- 1420070400000 0 IN IP4 127.0.0.1\r\ns=-\r\nt=0 0\r\na=msid-semantic: WMS *\r\na=group:BUNDLE 0 1\r\n\
369            m=audio {port} UDP/TLS/RTP/SAVPF {audio_payload}\r\na=rtpmap:{audio_payload} {audio_codec}/48000/2\r\na=fmtp:{audio_payload} minptime=10;useinbandfec=1;usedtx=0\r\na=rtcp-fb:{audio_payload} transport-cc\r\na=extmap:1 urn:ietf:params:rtp-hdrext:ssrc-audio-level\r\na=extmap:3 http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01\r\na=setup:{setup}\r\na=mid:{audio_mid}\r\na=maxptime:60\r\na={direction}\r\na=rtcp-mux\r\n\
370            m=video {port} UDP/TLS/RTP/SAVPF {video_payload} {video_rtxpayload}\r\na=rtpmap:{video_payload} {video_codec}/90000\r\na=rtpmap:{video_rtxpayload} rtx/90000\r\na=fmtp:{video_payload} x-google-max-bitrate=2500;level-asymmetry-allowed=1;packetization-mode=1;profile-level-id=42e01f\r\na=fmtp:{video_rtxpayload} apt={video_payload}\r\na=rtcp-fb:{video_payload} ccm fir\r\na=rtcp-fb:{video_payload} nack\r\na=rtcp-fb:{video_payload} nack pli\r\na=rtcp-fb:{video_payload} goog-remb\r\na=rtcp-fb:{video_payload} transport-cc\r\na=extmap:2 http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time\r\na=extmap:3 http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01\r\na=extmap:14 urn:ietf:params:rtp-hdrext:toffset\r\na=extmap:13 urn:3gpp:video-orientation\r\na=extmap:5 http://www.webrtc.org/experiments/rtp-hdrext/playout-delay\r\na=setup:{setup}\r\na=mid:{video_mid}\r\na={direction}\r\na=rtcp-mux\r\n"
371        );
372        answer.sdp = remote_sdp;
373
374        let mut parsed = answer.unmarshal()?;
375        for media in &mut parsed.media_descriptions {
376            media.connection_information = connection.clone();
377            for attribute in attributes {
378                media.attributes.push(attribute.clone());
379            }
380        }
381        let remote_sdp = parsed.marshal();
382        let recv_sdp = RTCSessionDescription::answer(remote_sdp)?;
383
384        peer_connection
385            .set_remote_description(recv_sdp.clone())
386            .await?;
387        debug!("[WebRTC] answer received, wait for quit event");
388        trace_tx
389            .as_ref()
390            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointRTCConnecting));
391        connected_rx.await?;
392
393        let local_audio_track = Arc::new(TrackLocalStaticSample::new(
394            RTCRtpCodecCapability {
395                mime_type: MIME_TYPE_OPUS.to_owned(),
396                ..Default::default()
397            },
398            "audio".to_owned(),
399            "webrtc-rs".to_owned(),
400        ));
401        audio_rtp_sender
402            .replace_track(Some(
403                Arc::clone(&local_audio_track) as Arc<dyn TrackLocal + Send + Sync>
404            ))
405            .await?;
406
407        let local_video_track = Arc::new(TrackLocalStaticSample::new(
408            RTCRtpCodecCapability {
409                mime_type: MIME_TYPE_H264.to_owned(),
410                ..Default::default()
411            },
412            "video".to_owned(),
413            "webrtc-rs".to_owned(),
414        ));
415        video_rtp_sender
416            .replace_track(Some(
417                Arc::clone(&local_video_track) as Arc<dyn TrackLocal + Send + Sync>
418            ))
419            .await?;
420
421        let user_id = user_id.get();
422        let channel_id = channel_id?;
423        dave_tx.send(DAVEPayload::OpCode4(
424            dave_protocol_version,
425            user_id,
426            channel_id,
427            local_audio_track,
428            local_video_track,
429        ))?;
430        dave_tx.send(DAVEPayload::Binary(external_payload))?;
431        trace_tx
432            .as_ref()
433            .map(|tx| tx.send(DiscordLiveBuilderState::EndpointDAVECreating));
434        let dave_instance = instance_rx.await?;
435
436        let payload = json!({
437            "op": 5,
438            "d": {
439                "speaking": 1,
440                "delay": 5,
441                "ssrc": 0
442            }
443        });
444        egress_tx.send(WebSocketMessage::text(payload.to_string()))?;
445
446        let payload = json!({
447            "op": 12,
448            "d": {
449                "audio_ssrc": audio_ssrc,
450                "video_ssrc": video_ssrc,
451                "rtx_ssrc": video_rtxssrc,
452                "streams": [{
453                    "type": "video",
454                    "rid": "100",
455                    "ssrc": video_ssrc,
456                    "active": true,
457                    "quality": 100,
458                    "rtx_ssrc": video_rtxssrc,
459                    "max_bitrate": 3500000,
460                    "max_framerate": 30,
461                    "max_resolution": {
462                        "type": "fixed",
463                        "width": 1280,
464                        "height": 720
465                    }
466                }]
467            }
468        });
469        let active = payload.to_string();
470        let payload = json!({
471            "op": 12,
472            "d": {
473                "audio_ssrc": 0,
474                "video_ssrc": streams[0].ssrc,
475                "rtx_ssrc": streams[0].rtx_ssrc,
476                "streams": [{
477                    "type": "video",
478                    "rid": "100",
479                    "ssrc": streams[0].ssrc,
480                    "active": false,
481                    "quality": 100,
482                    "rtx_ssrc": streams[0].rtx_ssrc,
483                    "max_bitrate": 3500000,
484                    "max_framerate": 30,
485                    "max_resolution": {
486                        "type": "fixed",
487                        "width": 1280,
488                        "height": 720
489                    }
490                }]
491            }
492        });
493        let inactive = payload.to_string();
494        egress_tx.send(WebSocketMessage::text(inactive))?;
495
496        let instance_lock = dave_instance.clone();
497        tokio::spawn(async move {
498            loop {
499                sleep(Duration::from_secs(300)).await;
500
501                let Ok(_) = peer_connection
502                    .set_remote_description(inactive_sdp.clone())
503                    .await
504                else {
505                    break;
506                };
507
508                let local_audio_track = Arc::new(TrackLocalStaticSample::new(
509                    RTCRtpCodecCapability {
510                        mime_type: MIME_TYPE_OPUS.to_owned(),
511                        ..Default::default()
512                    },
513                    "audio".to_owned(),
514                    "webrtc-rs".to_owned(),
515                ));
516                let Ok(_) = audio_rtp_sender
517                    .replace_track(Some(
518                        Arc::clone(&local_audio_track) as Arc<dyn TrackLocal + Send + Sync>
519                    ))
520                    .await
521                else {
522                    break;
523                };
524                {
525                    instance_lock
526                        .write()
527                        .await
528                        .replace_audio_track(local_audio_track);
529                }
530
531                let local_video_track = Arc::new(TrackLocalStaticSample::new(
532                    RTCRtpCodecCapability {
533                        mime_type: MIME_TYPE_H264.to_owned(),
534                        ..Default::default()
535                    },
536                    "video".to_owned(),
537                    "webrtc-rs".to_owned(),
538                ));
539                let Ok(_) = video_rtp_sender
540                    .replace_track(Some(
541                        Arc::clone(&local_video_track) as Arc<dyn TrackLocal + Send + Sync>
542                    ))
543                    .await
544                else {
545                    break;
546                };
547                {
548                    instance_lock
549                        .write()
550                        .await
551                        .replace_video_track(local_video_track);
552                }
553
554                let Ok(_) = peer_connection
555                    .set_remote_description(recv_sdp.clone())
556                    .await
557                else {
558                    break;
559                };
560            }
561        });
562
563        Ok(DiscordLive {
564            notify,
565            active,
566            dave_instance,
567            egress_tx,
568        })
569    }
570}
571
572pub struct DiscordLive {
573    notify: Arc<Notifier>,
574    active: String,
575    dave_instance: Arc<RwLock<DAVEInstance>>,
576    egress_tx: mpsc::UnboundedSender<WebSocketMessage>,
577}
578
579impl Mirror for DiscordLive {
580    fn write_audio_sample<'a>(
581        &'a self,
582        payload: &'a mut Sample,
583    ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + 'a>> {
584        Box::pin(async {
585            if self.notify.is_closed() {
586                return Err(Error {
587                    kind: ErrorType::DiscordEndpoint,
588                    source: None,
589                });
590            }
591            self.dave_instance
592                .write()
593                .await
594                .write_audio_sample(payload)
595                .await
596                .map_err(|err| Error {
597                    kind: ErrorType::DiscordEndpoint,
598                    source: Some(err.into()),
599                })
600        })
601    }
602
603    fn write_video_sample<'a>(
604        &'a self,
605        payload: &'a mut Sample,
606    ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + 'a>> {
607        Box::pin(async {
608            if self.notify.is_closed() {
609                return Err(Error {
610                    kind: ErrorType::DiscordEndpoint,
611                    source: None,
612                });
613            }
614            self.dave_instance
615                .write()
616                .await
617                .write_video_sample(payload)
618                .await
619                .map_err(|err| Error {
620                    kind: ErrorType::DiscordEndpoint,
621                    source: Some(err.into()),
622                })
623        })
624    }
625
626    fn call_connected_callback(&self) -> Result<(), Error> {
627        if self.notify.is_closed() {
628            return Err(Error {
629                kind: ErrorType::DiscordEndpoint,
630                source: None,
631            });
632        }
633        self.egress_tx
634            .send(WebSocketMessage::text(self.active.clone()))
635            .map_err(|err| Error {
636                kind: ErrorType::DiscordEndpoint,
637                source: Some(err.into()),
638            })
639    }
640
641    fn close(&self) {
642        self.notify.close()
643    }
644}
645
646struct DAVEInstance {
647    session: DaveSession,
648    dave_protocol_version: u16,
649    local_audio_track: Arc<TrackLocalStaticSample>,
650    local_video_track: Arc<TrackLocalStaticSample>,
651}
652
653impl DAVEInstance {
654    fn get_session(&mut self) -> &mut DaveSession {
655        &mut self.session
656    }
657
658    fn set_dave_protocol_version(&mut self, version: u16) -> u16 {
659        self.dave_protocol_version = version;
660        self.dave_protocol_version
661    }
662
663    fn replace_audio_track(&mut self, track: Arc<TrackLocalStaticSample>) {
664        self.local_audio_track = track;
665    }
666
667    fn replace_video_track(&mut self, track: Arc<TrackLocalStaticSample>) {
668        self.local_video_track = track;
669    }
670
671    async fn write_audio_sample(&mut self, payload: &mut Sample) -> Result<(), webrtc::Error> {
672        if self.dave_protocol_version == 0 || !self.session.is_ready() {
673            return self.local_audio_track.write_sample(payload).await;
674        }
675
676        let Ok(data) = self
677            .session
678            .encrypt(MediaType::AUDIO, Codec::OPUS, &payload.data)
679        else {
680            return self.local_audio_track.write_sample(payload).await;
681        };
682        payload.data = Bytes::copy_from_slice(&data);
683
684        self.local_audio_track.write_sample(payload).await
685    }
686
687    async fn write_video_sample(&mut self, payload: &mut Sample) -> Result<(), webrtc::Error> {
688        if self.dave_protocol_version == 0 || !self.session.is_ready() {
689            return self.local_video_track.write_sample(payload).await;
690        }
691
692        let mut data = Vec::new();
693        let mut nalu_indexes = Vec::new();
694        let mut i = 0;
695        while i < (payload.data.len() - NALU_SHORT_START_SEQUENCE_SIZE) {
696            if payload.data[i + 2] > START_CODE_HIGHEST_POSSIBLE_VALUE {
697                i += NALU_SHORT_START_SEQUENCE_SIZE;
698            } else if payload.data[i + 1] != START_CODE_LEADING_BYTES_VALUE {
699                i += 2;
700            } else if payload.data[i] != START_CODE_LEADING_BYTES_VALUE
701                || payload.data[i + 2] != START_CODE_END_BYTE_VALUE
702            {
703                i += 1;
704            } else {
705                if i >= 1 && payload.data[i - 1] == START_CODE_LEADING_BYTES_VALUE {
706                    nalu_indexes.push((i - 1, 4));
707                } else {
708                    nalu_indexes.push((i, 3));
709                }
710                i += NALU_SHORT_START_SEQUENCE_SIZE;
711            }
712        }
713
714        for pos in 0..nalu_indexes.len() {
715            let (nalu, start_size) = nalu_indexes[pos];
716            let next_nalu = nalu_indexes
717                .get(pos + 1)
718                .map(|v| v.0)
719                .unwrap_or(payload.data.len());
720            match payload.data[nalu + start_size] & 0x1F {
721                1 | 5 | 8 => {
722                    data.extend_from_slice(&payload.data[nalu..next_nalu]);
723                }
724                7 => {
725                    let (mut sps, _) =
726                        parse_sps(&payload.data[(nalu + start_size + 1)..next_nalu]).unwrap();
727                    if !sps.vui_parameters.bitstream_restriction_flag {
728                        sps.vui_parameters.bitstream_restriction_flag = true;
729                        sps.vui_parameters.motion_vectors_over_pic_boundaries_flag = true;
730                        sps.vui_parameters.max_bytes_per_pic_denom = 2;
731                        sps.vui_parameters.max_bits_per_mb_denom = 1;
732                        sps.vui_parameters.log2_max_mv_length_horizontal = 16;
733                        sps.vui_parameters.log2_max_mv_length_vertical = 16;
734                        sps.vui_parameters.max_num_reorder_frames = 0;
735                        sps.vui_parameters.max_dec_frame_buffering = sps.max_num_ref_frames as u32;
736                    }
737                    data.extend_from_slice(&payload.data[nalu..][..(start_size + 1)]);
738                    synthesize_sps(&sps, &mut data, false).unwrap();
739                }
740                _ => {}
741            }
742        }
743
744        let Ok(data) = self.session.encrypt(MediaType::VIDEO, Codec::H264, &data) else {
745            return self.local_video_track.write_sample(payload).await;
746        };
747        payload.data = Bytes::copy_from_slice(&data);
748
749        self.local_video_track.write_sample(payload).await
750    }
751}
752
753enum DAVEPayload {
754    Binary(Payload),
755    OpCode4(
756        u16,
757        u64,
758        u64,
759        Arc<TrackLocalStaticSample>,
760        Arc<TrackLocalStaticSample>,
761    ),
762    OpCode11(Vec<String>),
763    OpCode13(String),
764    OpCode21(u16, u16),
765    OpCode22(u16),
766    OpCode24(u16, u8),
767}
768
769pub(super) struct Notifier {
770    is_closed: AtomicBool,
771    gateway: Arc<Notify>,
772    endpoint: Arc<Notify>,
773    heartbeat: Arc<Notify>,
774    dave: Arc<Notify>,
775}
776
777impl Notifier {
778    fn new() -> Self {
779        Self {
780            is_closed: AtomicBool::new(false),
781            gateway: Arc::new(Notify::new()),
782            endpoint: Arc::new(Notify::new()),
783            heartbeat: Arc::new(Notify::new()),
784            dave: Arc::new(Notify::new()),
785        }
786    }
787
788    fn close(&self) {
789        self.gateway.notify_one();
790        self.endpoint.notify_one();
791        self.heartbeat.notify_one();
792        self.dave.notify_one();
793        self.is_closed.store(true, Ordering::Relaxed);
794    }
795
796    fn is_closed(&self) -> bool {
797        self.is_closed.load(Ordering::Relaxed)
798    }
799}
800
801pub enum DiscordLiveBuilderState {
802    VoiceConnecting,
803    StreamCreating,
804    EndpointWSConnecting,
805    EndpointWSSDP,
806    EndpointRTCCreating,
807    EndpointRTCNegotiation,
808    EndpointRTCConnecting,
809    EndpointDAVECreating,
810}
811
812impl Display for DiscordLiveBuilderState {
813    fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult {
814        match self {
815            DiscordLiveBuilderState::VoiceConnecting => f.write_str("connecting to voice channel"),
816            DiscordLiveBuilderState::StreamCreating => {
817                f.write_str("creating new live stream session")
818            }
819            DiscordLiveBuilderState::EndpointWSConnecting => {
820                f.write_str("connecting to live stream endpoint")
821            }
822            DiscordLiveBuilderState::EndpointWSSDP => {
823                f.write_str("waiting remote sdp from live stream endpoint")
824            }
825            DiscordLiveBuilderState::EndpointRTCCreating => f.write_str("creating new rtc client"),
826            DiscordLiveBuilderState::EndpointRTCNegotiation => {
827                f.write_str("rtc client currently applying all changes still pending")
828            }
829            DiscordLiveBuilderState::EndpointRTCConnecting => {
830                f.write_str("rtc client currently connecting to live stream endpoint")
831            }
832            DiscordLiveBuilderState::EndpointDAVECreating => {
833                f.write_str("creating new dave session")
834            }
835        }
836    }
837}
838
839pub trait ErrorInner: StdError + Send + Sync {}
840
841impl<T: StdError + Send + Sync> ErrorInner for T {}
842
843impl StdError for Error<dyn ErrorInner> {
844    fn source(&self) -> Option<&(dyn StdError + 'static)> {
845        self.source
846            .as_ref()
847            .map(|source| &**source as &(dyn StdError + 'static))
848    }
849}
850
851impl From<RecvError> for Error<dyn ErrorInner> {
852    fn from(err: RecvError) -> Self {
853        Self {
854            kind: ErrorType::DiscordIPC,
855            source: Some(Box::new(err)),
856        }
857    }
858}
859
860impl From<webrtc::Error> for Error<dyn ErrorInner> {
861    fn from(err: webrtc::Error) -> Self {
862        Self {
863            kind: ErrorType::DiscordEndpoint,
864            source: Some(Box::new(err)),
865        }
866    }
867}
868
869impl From<ParseIntError> for Error<dyn ErrorInner> {
870    fn from(err: ParseIntError) -> Self {
871        Self {
872            kind: ErrorType::DiscordEndpoint,
873            source: Some(Box::new(err)),
874        }
875    }
876}
877
878impl From<SendError<DAVEPayload>> for Error<dyn ErrorInner> {
879    fn from(err: SendError<DAVEPayload>) -> Self {
880        Self {
881            kind: ErrorType::DiscordIPC,
882            source: Some(Box::new(err)),
883        }
884    }
885}
886
887impl From<SendError<WebSocketMessage>> for Error<dyn ErrorInner> {
888    fn from(err: SendError<WebSocketMessage>) -> Self {
889        Self {
890            kind: ErrorType::DiscordIPC,
891            source: Some(Box::new(err)),
892        }
893    }
894}