use std::
{
time::Instant,
collections::HashMap,
sync::
{
Arc,
LazyLock,
Mutex as MutexSync,
},
};
use dashmap::DashMap;
use tokio::
{
task::AbortHandle,
net::tcp::OwnedWriteHalf,
sync::
{
Mutex,
mpsc::{ self, Sender },
},
};
use crate::
{
crypto,
consts::SharedKeys,
network::
{
self,
Streams,
codes::PacketCode,
server::
{
self,
connection::Connection,
},
screen::{ self, ScreenPacketCode },
},
};
static MUTED_ANIMATION: LazyLock<Vec<Vec<u8>>> = LazyLock::new(|| split_access_units(include_bytes!("./assets/muted.h264")));
static SHARES: LazyLock<DashMap<usize, Arc<Share>>> = LazyLock::new(|| DashMap::new());
struct ScreenTransferGuard
{
id: usize,
}
struct Share {
keyframe: MutexSync<Option<Vec<u8>>>, pending: MutexSync<Vec<(usize, Viewer)>>, }
struct Viewer {
token: [u8; 32], tx: Sender<ScreenPacketCode>, task: AbortHandle, needs_key: bool, }
impl Drop for ScreenTransferGuard
{
fn drop(&mut self)
{
SHARES.remove(&self.id);
if let Some(mut conn) = server::CONNECTIONS.iter_mut().find(|c| c.id() == Some(&self.id))
{
conn.remove_screen_stream();
}
}
}
impl Drop for Viewer
{
fn drop(&mut self)
{
self.task.abort();
}
}
fn split_access_units(bitstream: &[u8]) -> Vec<Vec<u8>> {
let mut units = Vec::new();
let mut start = None;
let mut index = 0;
while index + 3 < bitstream.len()
{
if bitstream[index..index + 3] != [0, 0, 1] { index += 1; continue; }
if bitstream[index + 3] & 0x1f == 7
{
let mut boundary = index;
while boundary > 0 && bitstream[boundary - 1] == 0 { boundary -= 1; }
if let Some(start) = start.replace(boundary)
{
units.push(bitstream[start..boundary].to_vec());
}
}
index += 3;
}
if let Some(start) = start { units.push(bitstream[start..].to_vec()); }
units
}
fn is_keyframe(bitstream: &[u8]) -> bool {
let mut index = 0;
while index + 3 < bitstream.len()
{
if bitstream[index..index + 3] != [0, 0, 1] { index += 1; continue; }
if matches!(bitstream[index + 3] & 0x1f, 5 | 7) { return true; }
index += 3;
}
false
}
fn spawn_viewer (
stream: Arc<Mutex<OwnedWriteHalf>>,
keys: &SharedKeys,
token: [u8; 32],
) -> Option<Viewer>
{
let mut rex_stream = crypto::init_rex_stream(keys, &token)?;
let (tx, mut rx) = mpsc::channel(screen::consts::VIEWER_CHANNEL_BOUND);
let task = tokio::spawn(async move
{
let mut seq = 0usize;
while let Some(code) = rx.recv().await
{
screen::send_frame(&mut *stream.lock().await, code, &mut rex_stream, Some(&mut seq)).await;
}
}).abort_handle();
Some(Viewer { token, tx, task, needs_key: true })
}
fn muted_frame(started: &Instant) -> Option<usize> {
let frames = MUTED_ANIMATION.len();
if frames == 0 { return None; }
Some((started.elapsed().as_millis() / screen::consts::MUTED_FRAME_INTERVAL.as_millis()) as usize % frames)
}
async fn end_share(id: usize) {
let (write_stream, keys, username) =
{
let mut conn = match server::CONNECTIONS.iter_mut().find(|c| c.id() == Some(&id))
{
Some(c) => c,
None => return
};
if conn.take_screen_stream().is_none() { return; }
(conn.write_stream().clone(), conn.keys().cloned(), conn.username().cloned())
};
log::info!("Screen share ended (upload socket closed): {}", server::log_addr(&id));
if let Some(username) = username
{
server::deattach(id, &username).await;
}
network::send(&mut *write_stream.lock().await, PacketCode::Screen { token: None }, keys.as_ref()).await;
}
pub fn attach (
sharer_id: usize,
client_id: usize,
keys: &SharedKeys,
stream: Arc<Mutex<OwnedWriteHalf>>,
token: [u8; 32],
) -> bool
{
let Some(share) = SHARES.get(&sharer_id).map(|share| share.clone()) else { return false; };
let Some(viewer) = spawn_viewer(stream, keys, token) else { return false; };
if let Some(frame) = share.keyframe.lock().ok().and_then(|frame| frame.clone())
{
let _ = viewer.tx.try_send(ScreenPacketCode::Video { data: frame });
}
match share.pending.lock()
{
Ok(mut pending) =>
{
pending.push((client_id, viewer));
true
},
Err(_) => false,
}
}
pub async fn screen(token: [u8; 32], id: usize, streams: &mut Streams<'_>, task: AbortHandle)
{
let keys =
{
let conn = server::CONNECTIONS.iter_mut()
.find(|e| e.value().id() == Some(&id));
match conn
{
Some(mut c) =>
{
let keys = match c.keys()
{
Some(k) => k.clone(),
None => return
};
c.set_screen_stream(task);
keys
},
None => return
}
};
let owner = server::log_addr(&id);
log::info!("Screen share started: {owner}");
let _guard = ScreenTransferGuard { id };
let mut seq = 0usize;
let mut viewers = HashMap::<usize, Viewer>::new();
let mut rex_stream = crypto::init_rex_stream(&keys, &token).unwrap();
let started = Instant::now();
let mut sent_muted_frame = None;
let share = Arc::new(Share
{
keyframe: MutexSync::new(None),
pending: MutexSync::new(Vec::new()),
});
SHARES.insert(id, share.clone());
loop
{
let read = match screen::receive_frame(streams, &mut rex_stream, &mut seq).await
{
Some(r) => r,
None => break
};
let muted = server::CONNECTIONS.iter()
.find(|c| c.id() == Some(&id))
.map(|c| *c.muted())
.unwrap_or(false);
let read = match (muted, read)
{
(false, read) =>
{
sent_muted_frame = None;
read
},
(true, ScreenPacketCode::Audio { .. }) => continue,
(true, ScreenPacketCode::Video { .. }) =>
{
let Some(frame) = muted_frame(&started) else { continue; };
if sent_muted_frame == Some(frame) { continue; }
sent_muted_frame = Some(frame);
ScreenPacketCode::Video { data: MUTED_ANIMATION[frame].clone() }
},
};
let arrivals = share.pending.lock().map(|mut pending| pending.drain(..).collect::<Vec<_>>()).unwrap_or_default();
for (client_id, viewer) in arrivals
{
viewers.insert(client_id, viewer);
log::info!("Screen viewer serving ({} attached): share of {owner}", viewers.len());
}
let entries: Vec<(usize, [u8; 32])> = server::CONNECTIONS.iter().filter_map(|entry|
{
match entry.value()
{
Connection::Authenticated { id: client_id, attached_screen, .. } =>
{
if let Some(attached_screen) = attached_screen && attached_screen.target_id == id
{
Some((*client_id, attached_screen.token))
} else { None }
},
_ => None,
}
}).collect();
viewers.retain(|client_id, viewer| entries.iter().any(|(e, token)| e == client_id && *token == viewer.token));
let keyframe = matches!(&read, ScreenPacketCode::Video { data } if is_keyframe(data));
for (client_id, viewer) in viewers.iter_mut()
{
if *client_id == id && matches!(read, ScreenPacketCode::Audio { .. }) { continue; }
if viewer.needs_key && matches!(read, ScreenPacketCode::Video { .. })
{
if !keyframe { continue; }
log::debug!("Screen viewer recovered on a keyframe: share of {owner}");
viewer.needs_key = false;
}
if viewer.tx.try_send(read.clone()).is_err()
{
if matches!(read, ScreenPacketCode::Video { .. })
{
if !viewer.needs_key { log::warn!("Screen viewer shed (link too slow): share of {owner}"); }
viewer.needs_key = true;
}
}
}
if keyframe && let ScreenPacketCode::Video { data } = &read
{
if let Ok(mut cached) = share.keyframe.lock() { *cached = Some(data.clone()); }
}
}
end_share(id).await;
}