diff --git a/virtio-devices/src/vsock/packet.rs b/virtio-devices/src/vsock/packet.rs index df6a1e7b5..4b2c60ad4 100644 --- a/virtio-devices/src/vsock/packet.rs +++ b/virtio-devices/src/vsock/packet.rs @@ -502,9 +502,7 @@ impl VsockPacket { /// Writes the local copy of the packet header to the guest memory. /// pub fn commit_hdr(&mut self, guest_mem: &M) -> Result<()> { - if self.len() as usize > defs::MAX_PKT_BUF_SIZE { - return Err(VsockError::InvalidPktLen(self.len())); - } + self.validate_len()?; guest_mem .write(self.hdr(), self.guest_hdr_addr) @@ -513,6 +511,19 @@ impl VsockPacket { Ok(()) } + fn validate_len(&self) -> Result<()> { + let len = self.len() as usize; + if len > defs::MAX_PKT_BUF_SIZE { + return Err(VsockError::InvalidPktLen(self.len())); + } + + match &self.buf { + Some(buf) if len > buf.len() => Err(VsockError::InvalidPktLen(self.len())), + None if len > 0 => Err(VsockError::PktBufMissing), + _ => Ok(()), + } + } + pub fn has_buf(&self) -> bool { self.buf.is_some() } @@ -1097,4 +1108,43 @@ mod unit_tests { .unwrap(); assert_eq!(&after, &payload); } + + #[test] + fn test_commit_hdr_allows_zero_length_packet() { + create_context!(test_ctx, handler_ctx); + let mut pkt = VsockPacket::from_rx_virtq_head( + &mut handler_ctx.handler.queues[0] + .iter(&test_ctx.mem) + .unwrap() + .next() + .unwrap(), + None, + ) + .unwrap(); + + assert_eq!(pkt.len(), 0); + pkt.commit_hdr(&test_ctx.mem).unwrap(); + } + + #[test] + fn test_commit_hdr_rejects_len_above_buf_capacity() { + create_context!(test_ctx, handler_ctx); + let mut pkt = VsockPacket::from_rx_virtq_head( + &mut handler_ctx.handler.queues[0] + .iter(&test_ctx.mem) + .unwrap() + .next() + .unwrap(), + None, + ) + .unwrap(); + + let cap = pkt.buf_capacity().unwrap() as u32; + pkt.set_len(cap + 1); + + match pkt.commit_hdr(&test_ctx.mem) { + Err(VsockError::InvalidPktLen(n)) => assert_eq!(n, cap + 1), + other => panic!("expected InvalidPktLen, got {other:?}"), + } + } }