Prevent infinite loops in device enumeration extensions #1483

This commit is contained in:
Andrew Case
2024-12-28 22:42:25 +00:00
parent b1a42d9510
commit 1bd031b9a8
@@ -405,11 +405,24 @@ class DEVICE_OBJECT(objects.StructType, pool.ExecutiveObject):
def get_attached_devices(self) -> Generator[ObjectInterface, None, None]:
"""Enumerate the attached device's objects"""
device = self.AttachedDevice.dereference()
while device:
yield device
device = device.AttachedDevice.dereference()
seen = set()
try:
device = self.AttachedDevice.dereference()
except exceptions.InvalidAddressException:
return
while device:
if device.vol.offset in seen:
break
seen.add(device.vol.offset)
yield device
try:
device = device.AttachedDevice.dereference()
except exceptions.InvalidAddressException:
return
class DRIVER_OBJECT(objects.StructType, pool.ExecutiveObject):
"""A class for kernel driver objects."""
@@ -421,10 +434,24 @@ class DRIVER_OBJECT(objects.StructType, pool.ExecutiveObject):
def get_devices(self) -> Generator[ObjectInterface, None, None]:
"""Enumerate the driver's device objects"""
device = self.DeviceObject.dereference()
seen = set()
try:
device = self.DeviceObject.dereference()
except exceptions.InvalidAddressException:
return
while device:
if device.vol.offset in seen:
return
seen.add(device.vol.offset)
yield device
device = device.NextDevice.dereference()
try:
device = device.NextDevice.dereference()
except exceptions.InvalidAddressException:
return
def is_valid(self) -> bool:
"""Determine if the object is valid."""