diff --git a/vmm/src/interrupt.rs b/vmm/src/interrupt.rs index e42ba2f76..70a58dfb1 100644 --- a/vmm/src/interrupt.rs +++ b/vmm/src/interrupt.rs @@ -5,7 +5,6 @@ use std::collections::HashMap; use std::io; -use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use devices::interrupt_controller::InterruptController; @@ -23,7 +22,7 @@ pub type Result = std::io::Result; struct InterruptRoute { gsi: u32, irq_fd: EventFd, - registered: AtomicBool, + registered: bool, } impl InterruptRoute { @@ -36,39 +35,39 @@ impl InterruptRoute { Ok(InterruptRoute { gsi, irq_fd, - registered: AtomicBool::new(false), + registered: false, }) } - pub fn enable(&self, vm: &dyn hypervisor::Vm) -> Result<()> { - if !self.registered.load(Ordering::Acquire) { + pub fn enable(&mut self, vm: &dyn hypervisor::Vm) -> Result<()> { + if !self.registered { vm.register_irqfd(&self.irq_fd, self.gsi) .map_err(|e| io::Error::other(format!("Failed registering irq_fd: {e}")))?; // Update internals to track the irq_fd as "registered". - self.registered.store(true, Ordering::Release); + self.registered = true; } Ok(()) } - pub fn disable(&self, vm: &dyn hypervisor::Vm) -> Result<()> { - if self.registered.load(Ordering::Acquire) { + pub fn disable(&mut self, vm: &dyn hypervisor::Vm) -> Result<()> { + if self.registered { vm.unregister_irqfd(&self.irq_fd, self.gsi) .map_err(|e| io::Error::other(format!("Failed unregistering irq_fd: {e}")))?; // Update internals to track the irq_fd as "unregistered". - self.registered.store(false, Ordering::Release); + self.registered = false; } Ok(()) } - pub fn trigger(&self) -> Result<()> { + pub fn trigger(&mut self) -> Result<()> { self.irq_fd.write(1) } - pub fn notifier(&self) -> Option { + pub fn notifier(&mut self) -> Option { Some( self.irq_fd .try_clone() @@ -85,7 +84,7 @@ pub struct RoutingEntry { pub struct MsiInterruptGroup { vm: Arc, gsi_msi_routes: Arc>>, - irq_routes: HashMap, + irq_routes: HashMap>, } impl MsiInterruptGroup { @@ -109,7 +108,7 @@ impl MsiInterruptGroup { fn new( vm: Arc, gsi_msi_routes: Arc>>, - irq_routes: HashMap, + irq_routes: HashMap>, ) -> Self { MsiInterruptGroup { vm, @@ -122,7 +121,7 @@ impl MsiInterruptGroup { impl InterruptSourceGroup for MsiInterruptGroup { fn enable(&self) -> Result<()> { for (_, route) in self.irq_routes.iter() { - route.enable(self.vm.as_ref())?; + route.lock().unwrap().enable(self.vm.as_ref())?; } Ok(()) @@ -130,7 +129,7 @@ impl InterruptSourceGroup for MsiInterruptGroup { fn disable(&self) -> Result<()> { for (_, route) in self.irq_routes.iter() { - route.disable(self.vm.as_ref())?; + route.lock().unwrap().disable(self.vm.as_ref())?; } Ok(()) @@ -138,7 +137,7 @@ impl InterruptSourceGroup for MsiInterruptGroup { fn trigger(&self, index: InterruptIndex) -> Result<()> { if let Some(route) = self.irq_routes.get(&index) { - return route.trigger(); + return route.lock().unwrap().trigger(); } Err(io::Error::other(format!( @@ -148,7 +147,7 @@ impl InterruptSourceGroup for MsiInterruptGroup { fn notifier(&self, index: InterruptIndex) -> Option { if let Some(route) = self.irq_routes.get(&index) { - return route.notifier(); + return route.lock().unwrap().notifier(); } None @@ -162,6 +161,7 @@ impl InterruptSourceGroup for MsiInterruptGroup { set_gsi: bool, ) -> Result<()> { if let Some(route) = self.irq_routes.get(&index) { + let mut route = route.lock().unwrap(); let entry = RoutingEntry { route: self.vm.make_routing_entry(route.gsi, &config), masked, @@ -293,10 +293,10 @@ impl InterruptManager for MsiInterruptManager { fn create_group(&self, config: Self::GroupConfig) -> Result> { let mut allocator = self.allocator.lock().unwrap(); - let mut irq_routes: HashMap = + let mut irq_routes: HashMap> = HashMap::with_capacity(config.count as usize); for i in config.base..config.base + config.count { - irq_routes.insert(i, InterruptRoute::new(&mut allocator)?); + irq_routes.insert(i, Mutex::new(InterruptRoute::new(&mut allocator)?)); } Ok(Arc::new(MsiInterruptGroup::new(