#![allow(unsafe_code)]
#![allow(clippy::redundant_pub_crate)]
use crate::{CaptureError, DeviceId};
use mediaway_common::{
Bytes, CodecKind, GpuBufferHandle, NativeHandle, PixelFormat, StreamInfo, VideoFrame,
VideoFrameStorage, VideoGeometry,
};
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock, PoisonError, Weak, mpsc};
use std::thread::JoinHandle;
use windows::Win32::Graphics::Direct3D11::{
D3D11_BIND_SHADER_RESOURCE, D3D11_TEXTURE2D_DESC, D3D11_USAGE_DEFAULT, ID3D11Device,
ID3D11DeviceContext, ID3D11Texture2D,
};
use windows::Win32::Graphics::Dxgi::Common::{DXGI_FORMAT_B8G8R8A8_UNORM, DXGI_SAMPLE_DESC};
use windows::Win32::Graphics::Dxgi::{
DXGI_ERROR_ACCESS_LOST, DXGI_OUTDUPL_FRAME_INFO, IDXGIDevice, IDXGIOutput1,
IDXGIOutputDuplication, IDXGIResource,
};
use windows::core::Interface;
const POLL_TIMEOUT_MS: u32 = 16;
const RING_DEPTH: usize = 3;
fn registry() -> &'static Mutex<HashMap<DeviceId, Weak<SharedDuplication>>> {
static REGISTRY: OnceLock<Mutex<HashMap<DeviceId, Weak<SharedDuplication>>>> = OnceLock::new();
REGISTRY.get_or_init(|| Mutex::new(HashMap::new()))
}
struct RingSlot {
raw_texture_ptr: usize,
generation: AtomicU64,
}
struct Ring {
slots: Vec<Arc<RingSlot>>,
latest: Mutex<Arc<RingSlot>>,
}
struct ConsumerRecord {
id: u64,
last_seen_generation: u64,
held: Option<Arc<RingSlot>>,
}
enum ControlMsg {
Attach {
reply: mpsc::Sender<Result<u64, CaptureError>>,
},
Detach {
id: u64,
},
}
pub(crate) struct SharedDuplication {
consumers: Arc<Mutex<Vec<ConsumerRecord>>>,
ring: Arc<Ring>,
control_tx: mpsc::Sender<ControlMsg>,
shutdown: Arc<AtomicBool>,
driver_thread: Mutex<Option<JoinHandle<()>>>,
stream_info: StreamInfo,
device_raw: usize,
}
impl Drop for SharedDuplication {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::SeqCst);
let handle = self
.driver_thread
.lock()
.unwrap_or_else(PoisonError::into_inner)
.take();
if let Some(handle) = handle {
let _ = handle.join();
}
}
}
pub(crate) fn attach(
key: DeviceId,
device_raw: usize,
output_index: u32,
) -> Result<(Arc<SharedDuplication>, u64, StreamInfo), CaptureError> {
let mut map = registry().lock().unwrap_or_else(PoisonError::into_inner);
let existing = map.get(&key).and_then(Weak::upgrade);
let shared = if let Some(shared) = existing {
shared
} else {
let shared = spawn_driver(device_raw, output_index)?;
map.insert(key, Arc::downgrade(&shared));
shared
};
if shared.device_raw != device_raw {
return Err(CaptureError::InvalidInput);
}
let (reply_tx, reply_rx) = mpsc::channel();
shared
.control_tx
.send(ControlMsg::Attach { reply: reply_tx })
.map_err(|_| CaptureError::Backend)?;
let consumer_id = reply_rx.recv().map_err(|_| CaptureError::Backend)??;
let stream_info = shared.stream_info.clone();
drop(map);
Ok((shared, consumer_id, stream_info))
}
fn spawn_driver(
device_raw: usize,
output_index: u32,
) -> Result<Arc<SharedDuplication>, CaptureError> {
let consumers: Arc<Mutex<Vec<ConsumerRecord>>> = Arc::new(Mutex::new(Vec::new()));
let shutdown = Arc::new(AtomicBool::new(false));
let (control_tx, control_rx) = mpsc::channel();
let (ready_tx, ready_rx) = mpsc::channel::<Result<(StreamInfo, Arc<Ring>), CaptureError>>();
let thread_consumers = Arc::clone(&consumers);
let thread_shutdown = Arc::clone(&shutdown);
let handle = std::thread::Builder::new()
.name("mediaway-dxgi-shared".to_owned())
.spawn(move || {
driver_loop(
device_raw,
output_index,
&thread_consumers,
&thread_shutdown,
&control_rx,
&ready_tx,
);
})
.map_err(|_| CaptureError::Backend)?;
let (stream_info, ring) = match ready_rx.recv() {
Ok(Ok(parts)) => parts,
Ok(Err(e)) => {
let _ = handle.join();
return Err(e);
}
Err(_) => {
let _ = handle.join();
return Err(CaptureError::Backend);
}
};
Ok(Arc::new(SharedDuplication {
consumers,
ring,
control_tx,
shutdown,
driver_thread: Mutex::new(Some(handle)),
stream_info,
device_raw,
}))
}
fn driver_loop(
device_raw: usize,
output_index: u32,
consumers: &Arc<Mutex<Vec<ConsumerRecord>>>,
shutdown: &Arc<AtomicBool>,
control_rx: &mpsc::Receiver<ControlMsg>,
ready_tx: &mpsc::Sender<Result<(StreamInfo, Arc<Ring>), CaptureError>>,
) {
let opened = open_duplication(device_raw, output_index);
let (device, duplication, stream_info) = match opened {
Ok(parts) => parts,
Err(e) => {
let _ = ready_tx.send(Err(e));
return;
}
};
let desc = texture_desc(&stream_info);
let ring_init = create_ring(&device, &desc);
let (ring, ring_textures) = match ring_init {
Ok(parts) => parts,
Err(e) => {
let _ = ready_tx.send(Err(e));
return;
}
};
let ring = Arc::new(ring);
if ready_tx.send(Ok((stream_info, Arc::clone(&ring)))).is_err() {
return;
}
let mut next_id: u64 = 0;
let mut transient: Vec<(Arc<RingSlot>, ID3D11Texture2D)> = Vec::new();
loop {
if shutdown.load(Ordering::SeqCst) {
break;
}
while let Ok(msg) = control_rx.try_recv() {
match msg {
ControlMsg::Attach { reply } => {
let id = attach_consumer(&mut next_id, consumers);
let _ = reply.send(Ok(id));
}
ControlMsg::Detach { id } => {
let mut guard = consumers.lock().unwrap_or_else(PoisonError::into_inner);
guard.retain(|c| c.id != id);
}
}
}
let mut frame_info = DXGI_OUTDUPL_FRAME_INFO::default();
let mut desktop_resource: Option<IDXGIResource> = None;
let acquire = unsafe {
duplication.AcquireNextFrame(
POLL_TIMEOUT_MS,
&raw mut frame_info,
&raw mut desktop_resource,
)
};
if let Err(e) = acquire {
if e.code() == DXGI_ERROR_ACCESS_LOST {
break; }
continue;
}
let Some(desktop_resource) = desktop_resource else {
continue;
};
let Ok(source_texture) = desktop_resource.cast::<ID3D11Texture2D>() else {
let _ = unsafe { duplication.ReleaseFrame() };
continue;
};
if let Ok(context) = unsafe { device.GetImmediateContext() } {
reclaim_transient(&mut transient);
publish_tick(
&device,
&context,
&source_texture,
&ring,
&ring_textures,
&mut transient,
&desc,
);
}
let _ = unsafe { duplication.ReleaseFrame() };
}
}
fn open_duplication(
device_raw: usize,
output_index: u32,
) -> Result<(ID3D11Device, IDXGIOutputDuplication, StreamInfo), CaptureError> {
let raw = device_raw as *mut std::ffi::c_void;
let device_ref =
unsafe { ID3D11Device::from_raw_borrowed(&raw) }.ok_or(CaptureError::InvalidInput)?;
let device: ID3D11Device = device_ref.clone();
let dxgi_device: IDXGIDevice = device.cast().map_err(|_| CaptureError::Backend)?;
let adapter = unsafe { dxgi_device.GetAdapter() }.map_err(|_| CaptureError::Backend)?;
let output =
unsafe { adapter.EnumOutputs(output_index) }.map_err(|_| CaptureError::InvalidInput)?;
let output1: IDXGIOutput1 = output.cast().map_err(|_| CaptureError::Backend)?;
let duplication =
unsafe { output1.DuplicateOutput(&device) }.map_err(|_| CaptureError::AccessDenied)?;
let dup_desc = unsafe { duplication.GetDesc() };
let width = dup_desc.ModeDesc.Width;
let height = dup_desc.ModeDesc.Height;
if width == 0 || height == 0 {
return Err(CaptureError::Backend);
}
let stream_info = StreamInfo::Video {
id: 0,
codec: CodecKind::RawVideo,
time_base: mediaway_common::Rational::new(1, 60),
geometry: VideoGeometry { width, height },
extra_data: Bytes::new(),
};
Ok((device, duplication, stream_info))
}
fn texture_desc(stream_info: &StreamInfo) -> D3D11_TEXTURE2D_DESC {
let geometry = stream_info.geometry().unwrap_or(VideoGeometry {
width: 0,
height: 0,
});
D3D11_TEXTURE2D_DESC {
Width: geometry.width,
Height: geometry.height,
MipLevels: 1,
ArraySize: 1,
Format: DXGI_FORMAT_B8G8R8A8_UNORM,
SampleDesc: DXGI_SAMPLE_DESC {
Count: 1,
Quality: 0,
},
Usage: D3D11_USAGE_DEFAULT,
BindFlags: D3D11_BIND_SHADER_RESOURCE.0 as u32,
MiscFlags: 0,
CPUAccessFlags: 0,
}
}
fn create_texture(
device: &ID3D11Device,
desc: &D3D11_TEXTURE2D_DESC,
) -> Result<ID3D11Texture2D, CaptureError> {
let mut texture: Option<ID3D11Texture2D> = None;
unsafe {
device
.CreateTexture2D(&raw const *desc, None, Some(&raw mut texture))
.map_err(|_| CaptureError::Backend)?;
}
texture.ok_or(CaptureError::Backend)
}
fn create_ring(
device: &ID3D11Device,
desc: &D3D11_TEXTURE2D_DESC,
) -> Result<(Ring, Vec<ID3D11Texture2D>), CaptureError> {
let mut slots = Vec::with_capacity(RING_DEPTH);
let mut textures = Vec::with_capacity(RING_DEPTH);
for _ in 0..RING_DEPTH {
let texture = create_texture(device, desc)?;
let raw_texture_ptr = Interface::as_raw(&texture) as usize;
slots.push(Arc::new(RingSlot {
raw_texture_ptr,
generation: AtomicU64::new(0),
}));
textures.push(texture);
}
let latest = Arc::clone(&slots[0]);
Ok((
Ring {
slots,
latest: Mutex::new(latest),
},
textures,
))
}
fn attach_consumer(next_id: &mut u64, consumers: &Arc<Mutex<Vec<ConsumerRecord>>>) -> u64 {
let id = *next_id;
*next_id += 1;
consumers
.lock()
.unwrap_or_else(PoisonError::into_inner)
.push(ConsumerRecord {
id,
last_seen_generation: 0,
held: None,
});
id
}
fn publish_tick(
device: &ID3D11Device,
context: &ID3D11DeviceContext,
source: &ID3D11Texture2D,
ring: &Ring,
ring_textures: &[ID3D11Texture2D],
transient: &mut Vec<(Arc<RingSlot>, ID3D11Texture2D)>,
desc: &D3D11_TEXTURE2D_DESC,
) {
let free = ring
.slots
.iter()
.position(|slot| Arc::strong_count(slot) == 1);
if let Some(index) = free {
let Some(dest_texture) = ring_textures.get(index) else {
return;
};
unsafe { context.CopyResource(dest_texture, source) };
let slot = &ring.slots[index];
slot.generation.fetch_add(1, Ordering::AcqRel);
*ring.latest.lock().unwrap_or_else(PoisonError::into_inner) = Arc::clone(slot);
return;
}
let Ok(texture) = create_texture(device, desc) else {
return; };
unsafe { context.CopyResource(&texture, source) };
let raw_texture_ptr = Interface::as_raw(&texture) as usize;
let slot = Arc::new(RingSlot {
raw_texture_ptr,
generation: AtomicU64::new(1),
});
*ring.latest.lock().unwrap_or_else(PoisonError::into_inner) = Arc::clone(&slot);
transient.push((slot, texture));
}
fn reclaim_transient(transient: &mut Vec<(Arc<RingSlot>, ID3D11Texture2D)>) {
transient.retain(|(slot, _texture)| Arc::strong_count(slot) > 1);
}
pub(crate) fn detach(shared: &SharedDuplication, consumer_id: u64) {
let _ = shared
.control_tx
.send(ControlMsg::Detach { id: consumer_id });
}
pub(crate) fn poll_shared_frame(
shared: &SharedDuplication,
consumer_id: u64,
next_pts: &mut i64,
) -> Result<Option<VideoFrame>, CaptureError> {
let mut guard = shared
.consumers
.lock()
.unwrap_or_else(PoisonError::into_inner);
let record = guard
.iter_mut()
.find(|c| c.id == consumer_id)
.ok_or(CaptureError::Closed)?;
if record.held.is_some() {
return Err(CaptureError::Backend);
}
let latest = Arc::clone(
&shared
.ring
.latest
.lock()
.unwrap_or_else(PoisonError::into_inner),
);
let generation = latest.generation.load(Ordering::Acquire);
if generation == 0 || generation == record.last_seen_generation {
return Ok(None);
}
let raw_texture_ptr = latest.raw_texture_ptr;
record.last_seen_generation = generation;
record.held = Some(latest);
drop(guard);
let geometry = shared.stream_info.geometry().unwrap_or(VideoGeometry {
width: 0,
height: 0,
});
let texture_handle = NativeHandle::new(raw_texture_ptr).ok_or(CaptureError::Backend)?;
let pts = *next_pts;
*next_pts += 1;
Ok(Some(VideoFrame {
pts,
duration: 1,
width: geometry.width,
height: geometry.height,
format: PixelFormat::Bgra8,
storage: VideoFrameStorage::Gpu(GpuBufferHandle::DirectX11 {
texture: texture_handle,
subresource: 0,
}),
}))
}
pub(crate) fn release_shared_frame(
shared: &SharedDuplication,
consumer_id: u64,
) -> Result<(), CaptureError> {
let mut guard = shared
.consumers
.lock()
.unwrap_or_else(PoisonError::into_inner);
let record = guard
.iter_mut()
.find(|c| c.id == consumer_id)
.ok_or(CaptureError::Closed)?;
record.held = None;
drop(guard);
Ok(())
}
#[cfg(test)]
#[path = "dxgi_shared_tests.rs"]
mod tests;