wdext 0.1.0

A DbgEng wrapper framework
// SPDX-FileCopyrightText: 2026 takubokudori
// SPDX-License-Identifier: MIT OR Apache-2.0
#[cfg(test)]
pub mod tests {
    use crate::tests::util::local_debug2;
    use std::{
        cell::{Cell, RefCell},
        rc::Rc,
    };
    use wdext::{
        DebuggeeOffset,
        callbacks::WdCallbacksErrorPolicy,
        data::{BreakpointId, DebugStatus},
        dbgeng::{
            DebugBreakpointRef, TimeoutQuery,
            callbacks::{
                DebugCesArg, DebugEventCallbacksHandler, DebugEventFlags,
            },
        },
    };

    mod util;

    #[derive(Debug, Clone, Copy, PartialEq, Eq)]
    enum ObservedEvent {
        ChangeEngineState,
        Breakpoint,
    }

    struct TestEventCallbacks {
        interest_mask_count: Cell<usize>,
        breakpoint_count: Cell<usize>,
        breakpoint_ids: RefCell<Vec<BreakpointId>>,
        breakpoint_offsets: RefCell<Vec<DebuggeeOffset>>,
        engine_state_changes: RefCell<Vec<DebugCesArg>>,
        call_order: RefCell<Vec<ObservedEvent>>,
    }

    impl TestEventCallbacks {
        fn new() -> Self {
            Self {
                interest_mask_count: Cell::new(0),
                breakpoint_count: Cell::new(0),
                breakpoint_ids: RefCell::new(Vec::new()),
                breakpoint_offsets: RefCell::new(Vec::new()),
                engine_state_changes: RefCell::new(Vec::new()),
                call_order: RefCell::new(Vec::new()),
            }
        }
    }

    impl DebugEventCallbacksHandler for TestEventCallbacks {
        fn get_interest_mask(&self) -> windows::core::Result<DebugEventFlags> {
            self.interest_mask_count
                .set(self.interest_mask_count.get() + 1);

            Ok(
                DebugEventFlags::Breakpoint
                    | DebugEventFlags::ChangeEngineState,
            )
        }

        fn breakpoint(&self, bp: DebugBreakpointRef) -> DebugStatus {
            self.breakpoint_count.set(self.breakpoint_count.get() + 1);

            self.breakpoint_ids.borrow_mut().push(bp.get_id().unwrap());

            self.breakpoint_offsets
                .borrow_mut()
                .push(bp.get_offset().unwrap());

            self.call_order.borrow_mut().push(ObservedEvent::Breakpoint);

            DebugStatus::Break
        }

        fn change_engine_state(
            &self,
            arg: DebugCesArg,
        ) -> windows::core::Result<()> {
            self.engine_state_changes.borrow_mut().push(arg);
            self.call_order
                .borrow_mut()
                .push(ObservedEvent::ChangeEngineState);
            Ok(())
        }
    }

    #[test]
    fn test_event_callbacks() {
        let ctx = local_debug2();
        let client = ctx.client();
        let control = ctx.control();
        let timeout = TimeoutQuery::from_secs(10);

        let main_offset = ctx
            .symbols()
            .get_symbol_offset_by_name("ListEntry!main")
            .unwrap();

        let callbacks = Rc::new(TestEventCallbacks::new());

        assert!(client.get_event_callbacks().unwrap().is_none());

        let _guard = client
            .set_forward_event_callbacks(
                callbacks.clone(),
                WdCallbacksErrorPolicy::ForwardError,
            )
            .unwrap();

        assert!(client.get_event_callbacks().unwrap().is_some());

        let breakpoint = control.add_software_breakpoint(main_offset).unwrap();

        let expected_breakpoint_id = breakpoint.get_id().unwrap();

        control.go(timeout).unwrap();

        assert!(callbacks.interest_mask_count.get() >= 1);

        assert_eq!(callbacks.breakpoint_count.get(), 1);

        assert_eq!(
            callbacks.breakpoint_ids.borrow().as_slice(),
            &[expected_breakpoint_id]
        );

        assert_eq!(
            callbacks.breakpoint_offsets.borrow().as_slice(),
            &[main_offset]
        );

        assert!(
            !callbacks.engine_state_changes.borrow().is_empty(),
            "ChangeEngineState was not forwarded to the handler"
        );

        assert!(
            callbacks
                .call_order
                .borrow()
                .contains(&ObservedEvent::Breakpoint)
        );

        breakpoint.remove().unwrap();

        client.clear_all_event_callbacks().unwrap();
        assert!(client.get_event_callbacks().unwrap().is_none());
    }
}