diff --git a/layersvt/device_memory_report/device_memory_report_handwritten_dispatch.cpp b/layersvt/device_memory_report/device_memory_report_handwritten_dispatch.cpp index 0c160a209d..aa7ca0b2fc 100644 --- a/layersvt/device_memory_report/device_memory_report_handwritten_dispatch.cpp +++ b/layersvt/device_memory_report/device_memory_report_handwritten_dispatch.cpp @@ -19,18 +19,22 @@ extern "C" { -static PFN_vkVoidFunction devmemreport_known_instance_functions(const char* pName) { +static PFN_vkVoidFunction devmemreport_known_global_functions(const char* pName) { if (strcmp(pName, "vkGetInstanceProcAddr") == 0) return reinterpret_cast(vkGetInstanceProcAddr); if (strcmp(pName, "vkCreateInstance") == 0) return reinterpret_cast(vkCreateInstance); + if (strcmp(pName, "vkEnumerateInstanceExtensionProperties") == 0) return reinterpret_cast(vkEnumerateInstanceExtensionProperties); + if (strcmp(pName, "vkEnumerateInstanceLayerProperties") == 0) return reinterpret_cast(vkEnumerateInstanceLayerProperties); + return nullptr; +} + +static PFN_vkVoidFunction devmemreport_known_instance_functions(const char* pName) { if (strcmp(pName, "vkDestroyInstance") == 0) return reinterpret_cast(vkDestroyInstance); if (strcmp(pName, "vkEnumeratePhysicalDevices") == 0) return reinterpret_cast(vkEnumeratePhysicalDevices); if (strcmp(pName, "vkEnumeratePhysicalDeviceGroups") == 0) return reinterpret_cast(vkEnumeratePhysicalDeviceGroups); - if (strcmp(pName, "vkEnumerateInstanceExtensionProperties") == 0) return reinterpret_cast(vkEnumerateInstanceExtensionProperties); - if (strcmp(pName, "vkEnumerateInstanceLayerProperties") == 0) return reinterpret_cast(vkEnumerateInstanceLayerProperties); return nullptr; } -static PFN_vkVoidFunction devmemreport_known_device_functions(const char* pName) { +static PFN_vkVoidFunction devmemreport_known_core_device_functions(const char* pName) { if (strcmp(pName, "vkGetDeviceProcAddr") == 0) return reinterpret_cast(vkGetDeviceProcAddr); if (strcmp(pName, "vkCreateDevice") == 0) return reinterpret_cast(vkCreateDevice); if (strcmp(pName, "vkDestroyDevice") == 0) return reinterpret_cast(vkDestroyDevice); @@ -41,29 +45,35 @@ static PFN_vkVoidFunction devmemreport_known_device_functions(const char* pName) if (strcmp(pName, "vkBindImageMemory") == 0) return reinterpret_cast(vkBindImageMemory); if (strcmp(pName, "vkBindBufferMemory2") == 0) return reinterpret_cast(vkBindBufferMemory2); if (strcmp(pName, "vkBindImageMemory2") == 0) return reinterpret_cast(vkBindImageMemory2); - if (strcmp(pName, "vkBindBufferMemory2KHR") == 0) return reinterpret_cast(vkBindBufferMemory2KHR); - if (strcmp(pName, "vkBindImageMemory2KHR") == 0) return reinterpret_cast(vkBindImageMemory2KHR); if (strcmp(pName, "vkCreateImage") == 0) return reinterpret_cast(vkCreateImage); if (strcmp(pName, "vkDestroyImage") == 0) return reinterpret_cast(vkDestroyImage); if (strcmp(pName, "vkCreateBuffer") == 0) return reinterpret_cast(vkCreateBuffer); if (strcmp(pName, "vkDestroyBuffer") == 0) return reinterpret_cast(vkDestroyBuffer); if (strcmp(pName, "vkGetImageMemoryRequirements") == 0) return reinterpret_cast(vkGetImageMemoryRequirements); if (strcmp(pName, "vkGetImageMemoryRequirements2") == 0) return reinterpret_cast(vkGetImageMemoryRequirements2); - if (strcmp(pName, "vkGetImageMemoryRequirements2KHR") == 0) return reinterpret_cast(vkGetImageMemoryRequirements2KHR); if (strcmp(pName, "vkGetBufferMemoryRequirements") == 0) return reinterpret_cast(vkGetBufferMemoryRequirements); if (strcmp(pName, "vkGetBufferMemoryRequirements2") == 0) return reinterpret_cast(vkGetBufferMemoryRequirements2); + return nullptr; +} + +static PFN_vkVoidFunction devmemreport_known_device_extension_functions(const char* pName) { + if (strcmp(pName, "vkBindBufferMemory2KHR") == 0) return reinterpret_cast(vkBindBufferMemory2KHR); + if (strcmp(pName, "vkBindImageMemory2KHR") == 0) return reinterpret_cast(vkBindImageMemory2KHR); + if (strcmp(pName, "vkGetImageMemoryRequirements2KHR") == 0) return reinterpret_cast(vkGetImageMemoryRequirements2KHR); if (strcmp(pName, "vkGetBufferMemoryRequirements2KHR") == 0) return reinterpret_cast(vkGetBufferMemoryRequirements2KHR); return nullptr; } -EXPORT_FUNCTION VKAPI_ATTR PFN_vkVoidFunction VKAPI_CALL vkGetInstanceProcAddr(VkInstance instance, const char* pName) { - PFN_vkVoidFunction func = devmemreport_known_instance_functions(pName); +static PFN_vkVoidFunction devmemreport_known_device_functions(const char* pName) { + PFN_vkVoidFunction func = devmemreport_known_core_device_functions(pName); if (func) { return func; } - - // If it's a device function, we can also return it here if we want to support GIPA for device functions. - func = devmemreport_known_device_functions(pName); + return devmemreport_known_device_extension_functions(pName); +} + +EXPORT_FUNCTION VKAPI_ATTR PFN_vkVoidFunction VKAPI_CALL vkGetInstanceProcAddr(VkInstance instance, const char* pName) { + PFN_vkVoidFunction func = devmemreport_known_global_functions(pName); if (func) { return func; } @@ -72,12 +82,34 @@ EXPORT_FUNCTION VKAPI_ATTR PFN_vkVoidFunction VKAPI_CALL vkGetInstanceProcAddr(V return nullptr; } + func = devmemreport_known_instance_functions(pName); + if (func) { + return func; + } + + // Core device functions can be returned directly from GIPA. + func = devmemreport_known_core_device_functions(pName); + if (func) { + return func; + } + auto table = instance_dispatch_table(instance); if (table == NULL || table->GetInstanceProcAddr == NULL) { return nullptr; } - return table->GetInstanceProcAddr(instance, pName); + // For extension device commands, verify the underlying chain supports them before returning an interceptor. + PFN_vkVoidFunction down_func = table->GetInstanceProcAddr(instance, pName); + if (down_func == nullptr) { + return nullptr; + } + + func = devmemreport_known_device_extension_functions(pName); + if (func) { + return func; + } + + return down_func; } EXPORT_FUNCTION VKAPI_ATTR PFN_vkVoidFunction VKAPI_CALL vkGetDeviceProcAddr(VkDevice device, const char* pName) { diff --git a/layersvt/device_memory_report/device_memory_report_handwritten_functions.h b/layersvt/device_memory_report/device_memory_report_handwritten_functions.h index 78dc41d8d4..6761dc3dbe 100644 --- a/layersvt/device_memory_report/device_memory_report_handwritten_functions.h +++ b/layersvt/device_memory_report/device_memory_report_handwritten_functions.h @@ -293,6 +293,7 @@ VKAPI_ATTR VkResult VKAPI_CALL vkBindBufferMemory2(VkDevice device, uint32_t bin // Intercept memory binding via vkBindBufferMemory2KHR to correlate buffer object handles with device memory allocations. VKAPI_ATTR VkResult VKAPI_CALL vkBindBufferMemory2KHR(VkDevice device, uint32_t bindInfoCount, const VkBindBufferMemoryInfo* pBindInfos) { + assert(device_dispatch_table(device)->BindBufferMemory2KHR != nullptr); VkResult result = device_dispatch_table(device)->BindBufferMemory2KHR(device, bindInfoCount, pBindInfos); if (result == VK_SUCCESS && pBindInfos != nullptr) { RecordBufferBindings(bindInfoCount, pBindInfos); @@ -319,6 +320,7 @@ VKAPI_ATTR VkResult VKAPI_CALL vkBindImageMemory2(VkDevice device, uint32_t bind // Intercept memory binding via vkBindImageMemory2KHR to correlate image object handles with device memory allocations. VKAPI_ATTR VkResult VKAPI_CALL vkBindImageMemory2KHR(VkDevice device, uint32_t bindInfoCount, const VkBindImageMemoryInfo* pBindInfos) { + assert(device_dispatch_table(device)->BindImageMemory2KHR != nullptr); VkResult result = device_dispatch_table(device)->BindImageMemory2KHR(device, bindInfoCount, pBindInfos); if (result == VK_SUCCESS && pBindInfos != nullptr) { RecordImageBinds(bindInfoCount, pBindInfos);