📄 src/face_detection.rs
use async_channel::{Receiver, RecvError, Sender, TryRecvError};
use bevy::{
    app::{App, Plugin, PreUpdate, Startup},
    ecs::{
        message::Message,
        resource::Resource,
        system::{Commands, ResMut},
    },
    log,
    math::Vec3,
    tasks::ComputeTaskPool,
};
use nokhwa::{
    Buffer, Camera, nokhwa_initialize,
    pixel_format::RgbFormat,
    utils::{CameraIndex, RequestedFormat, RequestedFormatType, Resolution},
};
use ort::{
    inputs,
    session::{Session, SessionOutputs},
    value::Tensor,
};

pub struct FaceDetectionPlugin;

impl Plugin for FaceDetectionPlugin {
    fn build(&self, app: &mut App) {
        app.add_message::<FaceMoved>()
            .add_systems(Startup, setup_face_detection)
            .add_systems(PreUpdate, publish_face_moved);
    }
}

#[derive(Message)]
pub struct FaceMoved(pub Vec3);

#[derive(Resource)]
struct FaceDetectionReader(Receiver<Vec3>);

fn setup_face_detection(mut commands: Commands) {
    let (frame_sender, frame_receiver) = async_channel::unbounded();
    let (face_sender, face_receiver) = async_channel::unbounded();

    ComputeTaskPool::get()
        .spawn(camera_loop(frame_sender))
        .detach();

    ComputeTaskPool::get()
        .spawn(detection_loop(frame_receiver, face_sender))
        .detach();

    commands.insert_resource(FaceDetectionReader(face_receiver));
}

async fn camera_loop(sender: Sender<Buffer>) {
    {
        let (init_sender, init_receiver) = async_channel::unbounded();
        nokhwa_initialize(move |success| {
            if let Err(err) = init_sender.send_blocking(success) {
                log::error!("Could not send nokhwa initialization ({success}): {err}");
            }
        });
        match init_receiver.recv().await {
            Ok(true) => {}
            Ok(false) => {
                log::error!("nokhwa failed to initialize");
                return;
            }
            Err(err) => {
                log::error!("nokhwa failed to initialize: {err}");
                return;
            }
        }
    }

    let format = RequestedFormat::new::<RgbFormat>(RequestedFormatType::None);
    let mut camera = match Camera::new(CameraIndex::Index(1), format) {
        Ok(camera) => camera,
        Err(err) => {
            log::error!("Failed to create camera: {err}");
            return;
        }
    };
    if let Err(err) = camera.open_stream() {
        log::error!("Failed to open camera stream: {err}");
        return;
    }

    loop {
        let frame = match camera.frame() {
            Ok(frame) => frame,
            Err(err) => {
                log::error!("Could not get frame: {err}");
                return;
            }
        };

        if let Err(err) = sender.send(frame).await {
            log::error!("Could not send frame to face thread: {err}");
            return;
        }
    }
}

async fn detection_loop(mut receiver: Receiver<Buffer>, sender: Sender<Vec3>) {
    let mut model = match Session::builder().and_then(|mut b| {
        b.commit_from_file("face_detection_yunet_2026may.onnx")
    }) {
        Ok(session) => session,
        Err(err) => {
            log::error!("Could not build ONNX session: {err}");
            return;
        }
    };

    loop {
        let frame = match recv_latest(&mut receiver).await {
            Ok(frame) => frame,
            Err(err) => {
                log::error!("Could not get frame from camera thread: {err}");
                return;
            }
        };
        let resolution = frame.resolution();
        let resolution = Resolution::new(
            resolution.width().div_ceil(32) * 32,
            resolution.height().div_ceil(32) * 32,
        );
        let image = match frame.decode_image::<RgbFormat>() {
            Ok(image) => image,
            Err(err) => {
                log::error!("Could not decode image: {err}");
                return;
            }
        };

        let mut tensor = match Tensor::from_array((
            [
                1,
                3,
                resolution.height() as usize,
                resolution.width() as usize,
            ],
            vec![0.0_f32; 1 * 3 * (resolution.height() * resolution.width()) as usize],
        )) {
            Ok(tensor) => tensor,
            Err(err) => {
                log::warn!("Could not create tensor: {err}");
                continue;
            }
        };

        for (y, row) in image.rows().enumerate() {
            for (x, pixel) in row.enumerate() {
                tensor[[0, 2, y as i64, x as i64]] = pixel[0] as f32;
                tensor[[0, 1, y as i64, x as i64]] = pixel[1] as f32;
                tensor[[0, 0, y as i64, x as i64]] = pixel[2] as f32;
            }
        }

        let outputs = match model.run(inputs!["input" => &tensor]) {
            Ok(output) => output,
            Err(err) => {
                log::error!("Model did not run to completion: {err}");
                return;
            }
        };

        let Some(face) = get_face_position(&outputs, &resolution, 32)
            .or_else(|| get_face_position(&outputs, &resolution, 16))
            .or_else(|| get_face_position(&outputs, &resolution, 8))
        else {
            log::info!("No face detected in frame");
            continue;
        };
        if let Err(err) = sender.send(face).await {
            log::warn!("Could not send face position: {err}");
        }
    }
}

const THRESHOLD: f32 = 0.75;
fn get_face_position(
    outputs: &SessionOutputs,
    resolution: &Resolution,
    stride: usize,
) -> Option<Vec3> {
    let cls = outputs.get(format!("cls_{stride}"))?;
    let obj = outputs.get(format!("obj_{stride}"))?;

    let (cls_shape, cls) = match cls.try_extract_tensor::<f32>() {
        Ok(cls) => cls,
        Err(err) => {
            log::warn!("cls data is not a float tensor: {err}");
            return None;
        }
    };
    let (obj_shape, obj) = match obj.try_extract_tensor::<f32>() {
        Ok(obj) => obj,
        Err(err) => {
            log::warn!("obj data is not a float tensor: {err}");
            return None;
        }
    };
    if cls_shape.len() != 3 || obj_shape.len() != 3 || cls_shape[1] != obj_shape[1] {
        log::warn!("Result cls and obj is on wrong format: {cls_shape} and {obj_shape}");
        return None;
    }

    let (max_i, max_score) = cls
        .iter()
        .zip(obj.iter())
        .enumerate()
        .map(|(i, (cls, obj))| (i, (cls.clamp(0.0, 1.0) * obj.clamp(0.0, 1.0)).sqrt()))
        .filter(|(_, s)| s.is_finite())
        .max_by(|(_, a), (_, b)| a.total_cmp(b))?;
    if max_score < THRESHOLD {
        log::info!("No face above threshold");
        return None;
    }

    let bbox = outputs.get(format!("bbox_{stride}"))?;
    let (bbox_shape, bbox) = match bbox.try_extract_tensor::<f32>() {
        Ok(bbox) => bbox,
        Err(err) => {
            log::warn!("bbox is not a float tensor: {err}");
            return None;
        }
    };
    if bbox_shape.len() != 3 || bbox_shape[1] <= max_i as i64 || bbox_shape[2] != 4 {
        log::warn!("Result bbox is on the wrong format: {bbox_shape}");
        return None;
    }

    let s = stride as f32;
    let cols = (resolution.width().div_ceil(32) * 32) as usize / stride;
    let (col, row) = ((max_i % cols) as f32, (max_i / cols) as f32);
    let x_pixels = (col + bbox[max_i * 4 + 0]) * s;
    let y_pixels = (row + bbox[max_i * 4 + 1]) * s;
    let width_pixels = bbox[max_i * 4 + 2].exp() * s;

    const HFOV_DEG: f32 = 60.0;
    const FACE_WIDTH_METERS: f32 = 0.2;

    let (w, h) = (resolution.width() as f32, resolution.height() as f32);
    let f = (w / 2.0) / (HFOV_DEG.to_radians() / 2.0).tan();
    let z = f * FACE_WIDTH_METERS / width_pixels;
    let pos = Vec3::new(
        (x_pixels - w / 2.0) * z / f,
        -(y_pixels - h / 2.0) * z / f,
        -z,
    );
    Some(pos)
}

fn publish_face_moved(
    mut commands: Commands,
    mut face_detection_reader: ResMut<FaceDetectionReader>,
) {
    let Ok(face) = try_recv_latest(&mut face_detection_reader.0) else {
        return;
    };

    commands.write_message(FaceMoved(face));
}

async fn recv_latest<T>(receiver: &mut Receiver<T>) -> Result<T, RecvError> {
    match try_recv_latest(receiver) {
        Ok(msg) => Ok(msg),
        Err(TryRecvError::Closed) => Err(RecvError),
        Err(TryRecvError::Empty) => receiver.recv().await,
    }
}

fn try_recv_latest<T>(receiver: &mut Receiver<T>) -> Result<T, TryRecvError> {
    receiver.try_recv().map(|mut msg| {
        while let Ok(next) = receiver.try_recv() {
            msg = next;
        }
        msg
    })
}