1use crate::{
9 arch::pci::{self, PciDevice},
10 hardware::virtio::{
11 common::{VirtioDevice, Virtqueue},
12 status,
13 },
14 memory,
15 sync::SpinLock,
16};
17use alloc::{boxed::Box, vec::Vec};
18use core::{mem, ptr, sync::atomic::Ordering};
19
20pub const SECTOR_SIZE: usize = 512;
22
23pub mod features {
25 pub const VIRTIO_BLK_F_SIZE_MAX: u32 = 1 << 1;
26 pub const VIRTIO_BLK_F_SEG_MAX: u32 = 1 << 2;
27 pub const VIRTIO_BLK_F_GEOMETRY: u32 = 1 << 4;
28 pub const VIRTIO_BLK_F_RO: u32 = 1 << 5;
29 pub const VIRTIO_BLK_F_BLK_SIZE: u32 = 1 << 6;
30 pub const VIRTIO_BLK_F_FLUSH: u32 = 1 << 9;
31 pub const VIRTIO_BLK_F_TOPOLOGY: u32 = 1 << 10;
32 pub const VIRTIO_BLK_F_CONFIG_WCE: u32 = 1 << 11;
33 pub const VIRTIO_BLK_F_DISCARD: u32 = 1 << 13;
34 pub const VIRTIO_BLK_F_WRITE_ZEROES: u32 = 1 << 14;
35}
36
37#[allow(dead_code)]
39#[repr(u32)]
40pub enum RequestType {
41 In = 0,
43 Out = 1,
45 Flush = 4,
47 GetId = 8,
49 Discard = 11,
51 WriteZeroes = 13,
53}
54
55#[repr(C)]
57#[derive(Debug, Clone, Copy)]
58pub struct BlockRequestHeader {
59 pub request_type: u32,
60 pub reserved: u32,
61 pub sector: u64,
62}
63
64#[repr(u8)]
66#[derive(Debug, Clone, Copy, PartialEq, Eq)]
67pub enum BlockStatus {
68 Ok = 0,
69 IoError = 1,
70 Unsupported = 2,
71}
72
73#[repr(C)]
75#[allow(dead_code)]
76struct BlockConfig {
77 capacity: u64,
78 size_max: u32,
79 seg_max: u32,
80 geometry_cylinders: u16,
81 geometry_heads: u8,
82 geometry_sectors: u8,
83 blk_size: u32,
84 }
86
87pub trait BlockDevice {
89 fn read_sector(&self, sector: u64, buf: &mut [u8]) -> Result<(), BlockError>;
91
92 fn write_sector(&self, sector: u64, buf: &[u8]) -> Result<(), BlockError>;
94
95 fn read_sectors(&self, sector: u64, count: u16, buf: &mut [u8]) -> Result<(), BlockError> {
101 let sector_size = SECTOR_SIZE;
102 for i in 0..count as u64 {
103 let off = (i as usize) * sector_size;
104 if off + sector_size > buf.len() {
105 return Err(BlockError::BufferTooSmall);
106 }
107 self.read_sector(sector + i, &mut buf[off..off + sector_size])?;
108 }
109 Ok(())
110 }
111
112 fn write_sectors(&self, sector: u64, count: u16, buf: &[u8]) -> Result<(), BlockError> {
117 let sector_size = SECTOR_SIZE;
118 for i in 0..count as u64 {
119 let off = (i as usize) * sector_size;
120 if off + sector_size > buf.len() {
121 return Err(BlockError::BufferTooSmall);
122 }
123 self.write_sector(sector + i, &buf[off..off + sector_size])?;
124 }
125 Ok(())
126 }
127
128 fn sector_count(&self) -> u64;
130}
131
132#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
133pub enum BlockError {
134 #[error("device I/O error")]
135 IoError,
136 #[error("invalid sector number")]
137 InvalidSector,
138 #[error("buffer too small")]
139 BufferTooSmall,
140 #[error("device not ready")]
141 NotReady,
142}
143
144const BOUNCE_POOL_SIZE: usize = 64 * 1024;
147
148struct BouncePool {
153 frame: memory::PhysFrame,
154 order: u8,
155}
156
157impl BouncePool {
158 unsafe fn allocate() -> Result<Self, BlockError> {
160 let pages = (BOUNCE_POOL_SIZE + 4095) / 4096;
161 let order = pages.next_power_of_two().trailing_zeros() as u8;
162 let frame =
163 crate::sync::with_irqs_disabled(|token| memory::allocate_phys_contiguous(token, order))
164 .map_err(|_| BlockError::NotReady)?;
165 Ok(Self { frame, order })
166 }
167
168 fn phys(&self) -> u64 {
169 self.frame.start_address.as_u64()
170 }
171
172 fn virt(&self) -> u64 {
173 crate::memory::phys_to_virt(self.phys())
174 }
175}
176
177impl Drop for BouncePool {
178 fn drop(&mut self) {
179 crate::sync::with_irqs_disabled(|token| {
180 memory::free_phys_contiguous(token, self.frame, self.order);
181 });
182 }
183}
184
185struct MetaPool {
188 frame: memory::PhysFrame,
189}
190
191impl MetaPool {
192 unsafe fn allocate() -> Result<Self, BlockError> {
193 let frame = crate::sync::with_irqs_disabled(|token| memory::allocate_frame(token))
194 .map_err(|_| BlockError::NotReady)?;
195 Ok(Self { frame })
196 }
197
198 fn phys(&self) -> u64 {
199 self.frame.start_address.as_u64()
200 }
201
202 fn virt(&self) -> u64 {
203 crate::memory::phys_to_virt(self.phys())
204 }
205
206 fn status_offset(&self) -> u64 {
208 mem::size_of::<BlockRequestHeader>() as u64
209 }
210}
211
212impl Drop for MetaPool {
213 fn drop(&mut self) {
214 crate::sync::with_irqs_disabled(|token| {
215 memory::free_frame(token, self.frame);
216 });
217 }
218}
219
220pub struct VirtioBlockDevice {
222 device: VirtioDevice,
223 queue: SpinLock<Virtqueue>,
224 capacity: u64,
225 block_size: u32,
226 bounce_pool: BouncePool,
228 meta_pool: MetaPool,
229}
230
231unsafe impl Send for VirtioBlockDevice {}
233unsafe impl Sync for VirtioBlockDevice {}
234
235static VIRTIO_BLK_WQ: crate::sync::WaitQueue = crate::sync::WaitQueue::new();
237
238static VIRTIO_BLK_DONE: core::sync::atomic::AtomicBool = core::sync::atomic::AtomicBool::new(false);
240
241static VIRTIO_BLK_ERROR: core::sync::atomic::AtomicBool =
243 core::sync::atomic::AtomicBool::new(false);
244
245impl VirtioBlockDevice {
246 pub unsafe fn new(pci_dev: PciDevice) -> Result<Self, &'static str> {
251 log::info!("VirtIO-blk: Initializing device at {:?}", pci_dev.address);
252
253 let bounce_pool = BouncePool::allocate().map_err(|_| "Failed to allocate bounce pool")?;
255 let meta_pool = MetaPool::allocate().map_err(|_| "Failed to allocate meta pool")?;
256
257 let device = VirtioDevice::new(pci_dev)?;
259
260 device.reset();
262
263 device.add_status(status::ACKNOWLEDGE as u8);
265
266 device.add_status(status::DRIVER as u8);
268
269 let device_features = device.read_device_features();
271 log::debug!("VirtIO-blk: Device features: 0x{:08x}", device_features);
272
273 let dev_feat = device_features;
278 let mut guest_features: u32 = 0;
279 let has_blk_size = dev_feat & (1 << 6) != 0;
280 if has_blk_size {
281 guest_features |= 1 << 6;
282 }
283 if dev_feat & (1 << 9) != 0 {
284 guest_features |= 1 << 9; }
286 if dev_feat & (1 << 29) != 0 {
287 guest_features |= 1 << 29; }
289 log::info!("VirtIO-blk: Negotiated features: 0x{:08x}", guest_features);
290 device.write_guest_features(guest_features);
291
292 device.add_status(status::FEATURES_OK as u8);
294
295 if device.get_status() & (status::FEATURES_OK as u8) == 0 {
297 return Err("Device doesn't support our feature set");
298 }
299
300 let queue_size = device.queue_max_size(0);
303 if queue_size == 0 {
304 return Err("VirtIO-blk queue 0 is unavailable");
305 }
306 log::info!("VirtIO-blk: queue 0 size = {}", queue_size);
307
308 let queue = Virtqueue::new(queue_size)?;
310
311 device.setup_queue(0, &queue);
313
314 device.add_status(status::DRIVER_OK as u8);
316
317 let capacity_low = device.read_reg_u32(20);
320 let capacity_high = device.read_reg_u32(24);
321 let capacity = ((capacity_high as u64) << 32) | (capacity_low as u64);
322
323 let blk_size = if has_blk_size {
325 let sz = device.read_reg_u32(28);
326 if sz == 0 {
327 SECTOR_SIZE as u32
328 } else {
329 sz
330 }
331 } else {
332 SECTOR_SIZE as u32
333 };
334
335 log::info!(
336 "VirtIO-blk: Capacity: {} sectors ({} MB), block_size={}",
337 capacity,
338 (capacity * SECTOR_SIZE as u64) / (1024 * 1024),
339 blk_size,
340 );
341
342 log::info!("VirtIO-blk: Device initialized successfully");
343
344 Ok(Self {
345 device,
346 queue: SpinLock::new(queue),
347 capacity,
348 block_size: blk_size,
349 bounce_pool,
350 meta_pool,
351 })
352 }
353
354 fn acquire_dma_buffer(
357 &self,
358 buf_size: usize,
359 is_write: bool,
360 src: Option<&[u8]>,
361 ) -> Result<(u64, u64, Option<(memory::PhysFrame, u8)>), BlockError> {
362 if buf_size <= BOUNCE_POOL_SIZE {
363 let phys = self.bounce_pool.phys();
364 let virt = self.bounce_pool.virt();
365 if is_write {
366 if let Some(s) = src {
367 unsafe {
368 ptr::copy_nonoverlapping(s.as_ptr(), virt as *mut u8, buf_size);
369 }
370 }
371 }
372 Ok((phys, virt, None))
373 } else {
374 let buf_pages = (buf_size + 4095) / 4096;
376 let buf_order = buf_pages.next_power_of_two().trailing_zeros() as u8;
377 let buf_frame = crate::sync::with_irqs_disabled(|token| {
378 memory::allocate_phys_contiguous(token, buf_order)
379 })
380 .map_err(|_| BlockError::NotReady)?;
381 let buf_phys = buf_frame.start_address.as_u64();
382 let buf_virt = crate::memory::phys_to_virt(buf_phys);
383 if is_write {
384 if let Some(s) = src {
385 unsafe {
386 ptr::copy_nonoverlapping(s.as_ptr(), buf_virt as *mut u8, buf_size);
387 }
388 }
389 }
390 Ok((buf_phys, buf_virt, Some((buf_frame, buf_order))))
391 }
392 }
393
394 fn release_dma_buffer(&self, allocated: Option<(memory::PhysFrame, u8)>) {
395 if let Some((frame, order)) = allocated {
396 crate::sync::with_irqs_disabled(|token| {
397 memory::free_phys_contiguous(token, frame, order);
398 });
399 }
400 }
401
402 fn do_request(
404 &self,
405 request_type: RequestType,
406 sector: u64,
407 mut data_buf: Option<(&mut [u8], bool)>, ) -> Result<(), BlockError> {
409 let meta_phys = self.meta_pool.phys();
411 let meta_virt = self.meta_pool.virt();
412 let status_off = self.meta_pool.status_offset();
413
414 let header_ptr = meta_virt as *mut BlockRequestHeader;
415 let status_ptr = (meta_virt + status_off) as *mut u8;
416 unsafe {
417 ptr::write(
418 header_ptr,
419 BlockRequestHeader {
420 request_type: request_type as u32,
421 reserved: 0,
422 sector,
423 },
424 );
425 ptr::write(status_ptr, 0xFF);
426 }
427
428 let mut data_alloc: Option<(memory::PhysFrame, u8)> = None;
430 let mut dma_buf_virt: u64 = 0;
431
432 let mut buffers = Vec::with_capacity(3);
433 buffers.push((
434 meta_phys,
435 mem::size_of::<BlockRequestHeader>() as u32,
436 false,
437 ));
438
439 if let Some((buf, is_write)) = data_buf.as_mut() {
440 let buf_size = buf.len();
441 let (dma_phys, dma_virt, alloc) =
442 self.acquire_dma_buffer(buf_size, *is_write, Some(buf))?;
443 data_alloc = alloc;
444 dma_buf_virt = dma_virt;
445
446 let device_writable = !*is_write;
447 buffers.push((dma_phys, buf_size as u32, device_writable));
448 }
449
450 buffers.push((meta_phys + status_off, 1, true));
452
453 let mut queue = self.queue.lock();
455 let token = match queue.add_buffer(&buffers) {
456 Ok(t) => t,
457 Err(e) => {
458 drop(queue);
459 self.release_dma_buffer(data_alloc);
460 log::error!("VirtIO-blk: add_buffer failed: {}", e);
461 return Err(BlockError::IoError);
462 }
463 };
464 if queue.should_notify() {
465 self.device.notify_queue(0);
466 }
467 drop(queue);
468
469 let has_data = data_buf.is_some();
471
472 if has_data && crate::process::current_task_id().is_some() {
473 VIRTIO_BLK_DONE.store(false, Ordering::Release);
475 VIRTIO_BLK_ERROR.store(false, Ordering::Release);
476 VIRTIO_BLK_WQ.wait_until(|| {
477 if VIRTIO_BLK_DONE.load(Ordering::Acquire) {
478 VIRTIO_BLK_DONE.store(false, Ordering::Release);
479 Some(())
480 } else {
481 None
482 }
483 });
484 if VIRTIO_BLK_ERROR.load(Ordering::Acquire) {
485 VIRTIO_BLK_ERROR.store(false, Ordering::Release);
486 self.release_dma_buffer(data_alloc);
487 return Err(BlockError::IoError);
488 }
489 } else {
490 let mut spins = 0u32;
492 loop {
493 let q = self.queue.lock();
494 if q.has_used() {
495 if let Some((t, _)) = q.peek_used() {
496 if t == token {
497 drop(q);
498 let mut q = self.queue.lock();
499 q.get_used();
500 break;
501 }
502 }
503 }
505 drop(q);
506 spins = spins.saturating_add(1);
507 if spins >= 5_000_000 {
508 let isr = self.device.read_isr_status();
509 log::error!(
510 "VirtIO-blk: timeout sector={} token={} isr={}",
511 sector,
512 token,
513 isr
514 );
515 self.release_dma_buffer(data_alloc);
516 return Err(BlockError::IoError);
517 }
518 core::hint::spin_loop();
519 }
520 }
521
522 let status_byte = unsafe { ptr::read(status_ptr) };
524
525 if let Some((buf, is_write)) = data_buf {
526 if !is_write && status_byte == BlockStatus::Ok as u8 {
527 unsafe {
528 ptr::copy_nonoverlapping(
529 dma_buf_virt as *const u8,
530 buf.as_mut_ptr(),
531 buf.len(),
532 );
533 }
534 }
535 }
536
537 self.release_dma_buffer(data_alloc);
538
539 if status_byte == BlockStatus::Ok as u8 {
540 Ok(())
541 } else {
542 log::error!("VirtIO-blk: Request failed with status {}", status_byte);
543 Err(BlockError::IoError)
544 }
545 }
546}
547
548impl BlockDevice for VirtioBlockDevice {
549 fn read_sector(&self, sector: u64, buf: &mut [u8]) -> Result<(), BlockError> {
551 if sector >= self.capacity {
552 return Err(BlockError::InvalidSector);
553 }
554 if buf.len() < SECTOR_SIZE {
555 return Err(BlockError::BufferTooSmall);
556 }
557 self.do_request(RequestType::In, sector, Some((buf, false)))
558 }
559
560 fn write_sector(&self, sector: u64, buf: &[u8]) -> Result<(), BlockError> {
562 if sector >= self.capacity {
563 return Err(BlockError::InvalidSector);
564 }
565 if buf.len() < SECTOR_SIZE {
566 return Err(BlockError::BufferTooSmall);
567 }
568 let buf_mut = buf.as_ptr() as *mut u8;
572 let buf_slice = unsafe { core::slice::from_raw_parts_mut(buf_mut, buf.len()) };
573 self.do_request(RequestType::Out, sector, Some((buf_slice, true)))
574 }
575
576 fn read_sectors(&self, sector: u64, count: u16, buf: &mut [u8]) -> Result<(), BlockError> {
582 let nbytes = (count as usize) * SECTOR_SIZE;
583 if sector.saturating_add(count as u64) > self.capacity {
584 return Err(BlockError::InvalidSector);
585 }
586 if buf.len() < nbytes {
587 return Err(BlockError::BufferTooSmall);
588 }
589 self.do_request(RequestType::In, sector, Some((buf, false)))
590 }
591
592 fn write_sectors(&self, sector: u64, count: u16, buf: &[u8]) -> Result<(), BlockError> {
594 let nbytes = (count as usize) * SECTOR_SIZE;
595 if sector.saturating_add(count as u64) > self.capacity {
596 return Err(BlockError::InvalidSector);
597 }
598 if buf.len() < nbytes {
599 return Err(BlockError::BufferTooSmall);
600 }
601 let buf_mut = buf.as_ptr() as *mut u8;
602 let buf_slice = unsafe { core::slice::from_raw_parts_mut(buf_mut, buf.len()) };
603 self.do_request(RequestType::Out, sector, Some((buf_slice, true)))
604 }
605
606 fn sector_count(&self) -> u64 {
608 self.capacity
609 }
610}
611
612static VIRTIO_BLOCK_PTR: core::sync::atomic::AtomicPtr<VirtioBlockDevice> =
614 core::sync::atomic::AtomicPtr::new(core::ptr::null_mut());
615
616static VIRTIO_BLOCK_IRQ: core::sync::atomic::AtomicU8 = core::sync::atomic::AtomicU8::new(0xFF);
618
619pub fn init() {
623 log::info!("VirtIO-blk: Scanning for devices...");
624
625 let pci_dev = match pci::probe_first(pci::ProbeCriteria {
628 vendor_id: Some(pci::vendor::VIRTIO),
629 device_id: Some(pci::device::VIRTIO_BLOCK),
630 class_code: Some(pci::class::MASS_STORAGE),
631 subclass: None,
632 prog_if: None,
633 })
634 .or_else(|| pci::find_virtio_device(pci::device::VIRTIO_BLOCK))
635 {
636 Some(dev) => dev,
637 None => {
638 log::warn!("VirtIO-blk: No block device found");
639 return;
640 }
641 };
642
643 let irq_line = pci_dev.read_config_u8(pci::config::INTERRUPT_LINE);
645
646 match unsafe { VirtioBlockDevice::new(pci_dev) } {
648 Ok(device) => {
649 let leaked: &'static mut VirtioBlockDevice = Box::leak(Box::new(device));
652 VIRTIO_BLOCK_PTR.store(leaked as *mut VirtioBlockDevice, Ordering::Release);
653 VIRTIO_BLOCK_IRQ.store(irq_line, Ordering::Relaxed);
654
655 crate::arch::idt::register_virtio_block_irq(irq_line);
657
658 log::info!("VirtIO-blk: Device initialized on IRQ {}", irq_line);
659 }
660 Err(e) => {
661 log::error!("VirtIO-blk: Failed to initialize device: {}", e);
662 }
663 }
664}
665
666pub fn handle_interrupt() {
671 let ptr = VIRTIO_BLOCK_PTR.load(Ordering::Acquire);
672 if ptr.is_null() {
673 return;
674 }
675
676 let device = unsafe { &*ptr };
678
679 let isr_status = device.device.read_isr_status();
680 if isr_status == 0 {
681 return; }
683
684 device.device.ack_interrupt();
686
687 VIRTIO_BLK_DONE.store(true, Ordering::Release);
692
693 VIRTIO_BLK_WQ.wake_one();
695
696 log::trace!("VirtIO-blk: Interrupt handled (ISR={})", isr_status);
697}
698
699pub fn get_device() -> Option<&'static VirtioBlockDevice> {
704 let ptr = VIRTIO_BLOCK_PTR.load(Ordering::Acquire);
705 if ptr.is_null() {
706 None
707 } else {
708 Some(unsafe { &*ptr })
710 }
711}
712
713pub fn get_irq() -> u8 {
715 VIRTIO_BLOCK_IRQ.load(Ordering::Relaxed)
716}