From dc924133c6155afa57f30558f1fbc6f37dde8810 Mon Sep 17 00:00:00 2001 From: Russell McGuire Date: Mon, 20 Jul 2026 15:53:04 -0700 Subject: [PATCH 1/2] Initial Feature commit for Extension Callback Tracing + Requires driver implementation for full support + Gated on existence of driver exporting the APIs needed Signed-off-by: Russell McGuire --- include/loader/ze_loader.h | 72 +++++++ source/drivers/null/ze_null.cpp | 101 +++++++++ source/drivers/null/ze_null.h | 53 ++++- source/lib/ze_lib.cpp | 90 +++++++- source/lib/ze_lib.h | 8 + source/loader/ze_loader.cpp | 26 +++ source/loader/ze_loader_internal.h | 6 + test/CMakeLists.txt | 24 ++- test/loader_ext_fn_callback.cpp | 316 +++++++++++++++++++++++++++++ 9 files changed, 692 insertions(+), 4 deletions(-) create mode 100644 test/loader_ext_fn_callback.cpp diff --git a/include/loader/ze_loader.h b/include/loader/ze_loader.h index f71f5392..cfc89372 100644 --- a/include/loader/ze_loader.h +++ b/include/loader/ze_loader.h @@ -565,6 +565,78 @@ zelDisableTracingLayer(void); ZE_DLLEXPORT ze_result_t ZE_APICALL zelGetTracingLayerState(bool* enabled); // Pointer to bool to receive tracing layer state +/////////////////////////////////////////////////////////////////////////////// +/// @brief Callback signature for extension-function prologue/epilogue handlers. +/// +/// This intentionally mirrors the established per-API tracing callback shape +/// (see the ze_pfnXCb_t typedefs in ze_api.h) so tools can reuse their existing +/// callback infrastructure. Because an arbitrary extension function has no +/// generated params struct, @p pParams is passed as an opaque void* whose layout +/// is defined by the driver for the named function (may be null for pure-vendor +/// functions). The identity of the fired function is carried via +/// @p pTracerUserData (set at registration time). +/// +/// @param[in] pParams driver-populated parameter block (opaque) +/// @param[in] result epilogue only: the function's return value +/// @param[in] pTracerUserData per-registration user data +/// @param[in,out] ppTracerInstanceUserData per-call scratch for prologue->epilogue handoff +typedef void (ZE_APICALL *zel_pfnDriverExtensionFunctionCb_t)( + void* pParams, + ze_result_t result, + void* pTracerUserData, + void** ppTracerInstanceUserData + ); + +/////////////////////////////////////////////////////////////////////////////// +/// @brief Signature of the per-driver hook that enables or disables the driver's +/// extension-function callbacks. +/// +/// A driver that supports extension-function tracing exposes this by name +/// ("zelDriverEnableTracing") via zeDriverGetExtensionFunctionAddress. The loader +/// calls it on each active driver when the tracing layer is enabled/disabled +/// (including static ZE_ENABLE_TRACING_LAYER enablement and late-loaded drivers). +/// When disabled, the driver must not invoke any registered prologue/epilogue. +typedef ze_result_t (ZE_APICALL *zel_pfnDriverEnableTracing_t)( + ze_driver_handle_t hDriver, + ze_bool_t enable + ); + +/////////////////////////////////////////////////////////////////////////////// +/// @brief Registers prologue/epilogue callbacks for a named extension function. +/// +/// Extension functions obtained by string name via +/// zeDriverGetExtensionFunctionAddress() return a raw driver pointer that the +/// application calls directly, bypassing the loader and therefore the tracing +/// layer. This API provides a driver-side interception hook: the specified +/// driver invokes the registered @p prologue before, and @p epilogue after, the +/// body of the extension function named @p functionName. +/// +/// Registration is keyed by @p functionName and is order-independent relative to +/// zeDriverGetExtensionFunctionAddress() — it takes effect on the next invocation +/// of the function, even if the application already cached the function pointer. +/// Passing null for both @p prologue and @p epilogue unregisters the callbacks. +/// +/// @param[in] hDriver handle of the driver instance +/// @param[in] functionName name of the extension function to intercept +/// @param[in] pUserData user data passed to the callbacks as pTracerUserData +/// @param[in] prologue handler invoked before the function body (may be null) +/// @param[in] epilogue handler invoked after the function body (may be null) +/// +/// @return +/// - ZE_RESULT_SUCCESS on success (including unregister). +/// - ZE_RESULT_ERROR_UNINITIALIZED if the loader is not initialized. +/// - ZE_RESULT_ERROR_UNSUPPORTED_FEATURE if the driver does not implement this hook. +/// - ZE_RESULT_ERROR_INVALID_NULL_HANDLE if @p hDriver is null. +/// - ZE_RESULT_ERROR_INVALID_NULL_POINTER if @p functionName is null. +ZE_DLLEXPORT ze_result_t ZE_APICALL +zelDriverSetExtensionFunctionCallback( + ze_driver_handle_t hDriver, // [in] handle of the driver instance + const char* functionName, // [in] extension function name to intercept + void* pUserData, // [in][optional] user data passed to callbacks + zel_pfnDriverExtensionFunctionCb_t prologue, // [in][optional] prologue handler + zel_pfnDriverExtensionFunctionCb_t epilogue // [in][optional] epilogue handler + ); + #if defined(__cplusplus) } // extern "C" #endif diff --git a/source/drivers/null/ze_null.cpp b/source/drivers/null/ze_null.cpp index 2418e033..832e8a86 100644 --- a/source/drivers/null/ze_null.cpp +++ b/source/drivers/null/ze_null.cpp @@ -47,6 +47,33 @@ namespace driver return ZE_RESULT_SUCCESS; }; + ////////////////////////////////////////////////////////////////////////// + // Custom extension-function resolver. Returns real driver pointers by name + // for the setter and the sample extension function (the generic intercept + // in ze_nullddi.cpp defers to this hook and forwards *ppFunctionAddress). + zeDdiTable.Driver.pfnGetExtensionFunctionAddress = []( + ze_driver_handle_t, + const char* name, + void** ppFunctionAddress ) + { + if( nullptr == name || nullptr == ppFunctionAddress ) + return ZE_RESULT_ERROR_INVALID_NULL_POINTER; + if( 0 == strcmp( name, "zelDriverSetExtensionFunctionCallback" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zelDriverSetExtensionFunctionCallback ); + return ZE_RESULT_SUCCESS; + } + if( 0 == strcmp( name, "zelDriverEnableTracing" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zelDriverEnableTracing ); + return ZE_RESULT_SUCCESS; + } + if( 0 == strcmp( name, "zeSampleExtFunc" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zeSampleExtFunc ); + return ZE_RESULT_SUCCESS; + } + *ppFunctionAddress = nullptr; + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + }; + ////////////////////////////////////////////////////////////////////////// zeDdiTable.Device.pfnGet = []( ze_driver_handle_t, @@ -680,6 +707,80 @@ namespace driver pRuntime.version = ZE_API_VERSION_CURRENT; } + /////////////////////////////////////////////////////////////////////////// + /// @brief Sample extension function reachable only by name. Its body invokes + /// any registered prologue/epilogue with a typed params block. + ze_result_t ZE_APICALL zeSampleExtFunc( + ze_driver_handle_t hDriver, uint32_t input, uint32_t* pOutput ) + { + // Snapshot any registered callbacks for this function (name-keyed). + context_t::extension_function_callbacks_t cbs; + bool haveCbs = false; + { + std::lock_guard lock( context.extensionCallbackMutex ); + auto it = context.extensionCallbacks.find( "zeSampleExtFunc" ); + if( it != context.extensionCallbacks.end() ) { + cbs = it->second; + haveCbs = true; + } + } + + // Two-level gate: callbacks fire only when tracing is globally enabled + // AND a callback is registered for this function. + const bool fire = haveCbs && context.extensionCallbacksEnabled.load(); + + // Typed parameter block the driver exposes to the callbacks. + ze_sample_ext_func_params_t params = { &hDriver, &input, &pOutput }; + void* pInstanceData = nullptr; + ze_result_t result = ZE_RESULT_SUCCESS; + + if( fire && nullptr != cbs.prologue ) + cbs.prologue( ¶ms, result, cbs.pUserData, &pInstanceData ); + + // The (trivial) work of the extension function. + if( nullptr != pOutput ) + *pOutput = input * 2; + + if( fire && nullptr != cbs.epilogue ) + cbs.epilogue( ¶ms, result, cbs.pUserData, &pInstanceData ); + + return result; + } + + /////////////////////////////////////////////////////////////////////////// + /// @brief Enable/disable this driver's extension-function callbacks (the + /// global gate). Called by the loader when the tracing layer is + /// enabled/disabled. + ze_result_t ZE_APICALL zelDriverEnableTracing( + ze_driver_handle_t /*hDriver*/, ze_bool_t enable ) + { + context.extensionCallbacksEnabled.store( enable != 0 ); + return ZE_RESULT_SUCCESS; + } + + /////////////////////////////////////////////////////////////////////////// + /// @brief Driver-side registration entry (resolved by name from the loader). + /// Permissive and name-keyed: any name registers; null+null unregisters. + ze_result_t ZE_APICALL zelDriverSetExtensionFunctionCallback( + ze_driver_handle_t, const char* functionName, void* pUserData, + zel_pfnDriverExtensionFunctionCb_t prologue, + zel_pfnDriverExtensionFunctionCb_t epilogue ) + { + if( nullptr == functionName ) + return ZE_RESULT_ERROR_INVALID_NULL_POINTER; + + std::lock_guard lock( context.extensionCallbackMutex ); + if( nullptr == prologue && nullptr == epilogue ) { + context.extensionCallbacks.erase( functionName ); + } else { + auto& entry = context.extensionCallbacks[ functionName ]; + entry.pUserData = pUserData; + entry.prologue = prologue; + entry.epilogue = epilogue; + } + return ZE_RESULT_SUCCESS; + } + char *context_t::setenv_var_with_driver_id(const std::string &key, uint32_t driverId) { std::string env = key + "=" + std::to_string(driverId); diff --git a/source/drivers/null/ze_null.h b/source/drivers/null/ze_null.h index afff2802..f0849b07 100644 --- a/source/drivers/null/ze_null.h +++ b/source/drivers/null/ze_null.h @@ -10,11 +10,16 @@ #pragma once #include #include +#include +#include +#include +#include #include "ze_ddi.h" #include "zet_ddi.h" #include "zes_ddi.h" #include "ze_util.h" #include "ze_ddi_common.h" +#include "loader/ze_loader.h" #ifndef ZEL_NULL_DRIVER_ID #define ZEL_NULL_DRIVER_ID 1 @@ -47,6 +52,23 @@ namespace driver std::vector globalBaseNullHandle; bool ddiExtensionSupported = false; std::vector env_vars{}; + + // Registry for zelDriverSetExtensionFunctionCallback: maps an extension + // function name to the prologue/epilogue the driver invokes from inside + // that function's body. Keyed by name (order-independent vs fetch). + struct extension_function_callbacks_t { + void* pUserData = nullptr; + zel_pfnDriverExtensionFunctionCb_t prologue = nullptr; + zel_pfnDriverExtensionFunctionCb_t epilogue = nullptr; + }; + std::mutex extensionCallbackMutex; + std::map extensionCallbacks; + + // Global gate for extension-function callbacks, toggled by the loader via + // zelDriverEnableTracing. Callbacks fire only when this is set AND a + // callback is registered for the function (two-level gate). + std::atomic extensionCallbacksEnabled{false}; + context_t(); ~context_t(); @@ -68,7 +90,36 @@ namespace driver uint32_t ZE_APICALL zerTranslateDeviceHandleToIdentifier(ze_device_handle_t hDevice); ze_device_handle_t ZE_APICALL zerTranslateIdentifierToDeviceHandle(uint32_t identifier); ze_context_handle_t ZE_APICALL zerGetDefaultContext(void); - + + /////////////////////////////////////////////////////////////////////////// + // Extension-function callback prototype demonstration. + // + // "zeSampleExtFunc" is a stand-in vendor extension function reachable only by + // name via zeDriverGetExtensionFunctionAddress. Its body invokes any + // prologue/epilogue registered through zelDriverSetExtensionFunctionCallback, + // passing a typed params block (the driver knows its own signature). + typedef struct _ze_sample_ext_func_params_t + { + ze_driver_handle_t* phDriver; + uint32_t* pinput; + uint32_t** ppOutput; + } ze_sample_ext_func_params_t; + + ze_result_t ZE_APICALL zeSampleExtFunc( + ze_driver_handle_t hDriver, uint32_t input, uint32_t* pOutput ); + + // Driver-side registration entry, resolved by name from the loader's + // zelDriverSetExtensionFunctionCallback forward. + ze_result_t ZE_APICALL zelDriverSetExtensionFunctionCallback( + ze_driver_handle_t hDriver, const char* functionName, void* pUserData, + zel_pfnDriverExtensionFunctionCb_t prologue, + zel_pfnDriverExtensionFunctionCb_t epilogue ); + + // Driver-side enable/disable of extension-function callbacks, resolved by + // name from the loader when the tracing layer is enabled/disabled. + ze_result_t ZE_APICALL zelDriverEnableTracing( + ze_driver_handle_t hDriver, ze_bool_t enable ); + extern context_t context; } // namespace driver diff --git a/source/lib/ze_lib.cpp b/source/lib/ze_lib.cpp index 14e20511..e2aec2af 100644 --- a/source/lib/ze_lib.cpp +++ b/source/lib/ze_lib.cpp @@ -205,6 +205,10 @@ namespace ze_lib if (loaderGetContext == nullptr) { std::string message = "ze_lib Context Init() zelLoaderGetContext missing"; debug_trace_message(message, ""); + } else { + // Cache the loader context portably (the loader symbols are not + // link-time visible in the static-loader build). + ze_lib::context->loaderContext = loaderGetContext(); } std::string version_message = "Loader API Version to be requested is v" + std::to_string(ZE_MAJOR_VERSION(version)) + "." + std::to_string(ZE_MINOR_VERSION(version)); @@ -215,6 +219,7 @@ namespace ze_lib if( ZE_RESULT_SUCCESS == result ) { tracing_lib = zeLoaderGetTracingHandle(); } + ze_lib::context->loaderContext = loader::context; #endif @@ -635,6 +640,12 @@ zelEnableTracingLayer() if (ze_lib::context->pTracingZerDdiTable != nullptr) { ze_lib::context->zerDdiTable.exchange(ze_lib::context->pTracingZerDdiTable); } + // Propagate the enable to each active driver's extension-function tracing. + if (loader::context) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, true); + } + } } #endif return ZE_RESULT_SUCCESS; @@ -684,14 +695,91 @@ zelDisableTracingLayer() if (ze_lib::context->dynamicTracingSupported == false) { return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; } - if (ze_lib::context->tracingLayerEnableCounter.fetch_sub(1) <= 1) { + // Guard against underflow: decrement only when the counter is > 0, so a + // disable with no matching enable (e.g. under ZE_ENABLE_TRACING_LAYER=1) is a + // safe no-op rather than wrapping the unsigned counter and corrupting state. + uint32_t prev = ze_lib::context->tracingLayerEnableCounter.load(); + while (prev > 0 && + !ze_lib::context->tracingLayerEnableCounter.compare_exchange_weak(prev, prev - 1)) { + // prev is reloaded by compare_exchange_weak on failure + } + if (prev == 1) { + // 1 -> 0 transition: tear down the dynamic tracing DDI tables. ze_lib::context->zeDdiTable.exchange(&ze_lib::context->initialzeDdiTable); if (ze_lib::context->pTracingZerDdiTable != nullptr) { ze_lib::context->zerDdiTable.exchange(&ze_lib::context->initialzerDdiTable); } + // Disable per-driver extension tracing, unless tracing was enabled + // statically via ZE_ENABLE_TRACING_LAYER (documented to stay on for the + // whole application) - respect that sticky state. + if (loader::context && !loader::context->tracingLayerEnabled) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, false); + } + } } #endif return ZE_RESULT_SUCCESS; } +ze_result_t ZE_APICALL +zelDriverSetExtensionFunctionCallback( + ze_driver_handle_t hDriver, + const char* functionName, + void* pUserData, + zel_pfnDriverExtensionFunctionCb_t prologue, + zel_pfnDriverExtensionFunctionCb_t epilogue + ) +{ + if( nullptr == hDriver ) + return ZE_RESULT_ERROR_INVALID_NULL_HANDLE; + if( nullptr == functionName ) + return ZE_RESULT_ERROR_INVALID_NULL_POINTER; + if( ze_lib::destruction ) + return ZE_RESULT_ERROR_UNINITIALIZED; + + // Type of the driver-side registration entry point, discovered by name. + typedef ze_result_t (ZE_APICALL *zelDriverSetExtensionFunctionCallback_t)( + ze_driver_handle_t, const char*, void*, + zel_pfnDriverExtensionFunctionCb_t, zel_pfnDriverExtensionFunctionCb_t ); + + // Resolve the driver's registration entry via the standard extension-address + // lookup on this specific driver. The driver owns the registry and the + // invocation of the callbacks from inside the extension-function body. + auto pfnGetExtensionFunctionAddress = + ze_lib::context->zeDdiTable.load()->Driver.pfnGetExtensionFunctionAddress; + if( nullptr == pfnGetExtensionFunctionAddress ) { + if( !ze_lib::context->isInitialized ) + return ZE_RESULT_ERROR_UNINITIALIZED; + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + } + + void* pfnRaw = nullptr; + ze_result_t result = pfnGetExtensionFunctionAddress( + hDriver, "zelDriverSetExtensionFunctionCallback", &pfnRaw ); + if( result != ZE_RESULT_SUCCESS ) + return result; + if( nullptr == pfnRaw ) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + + // Sync this driver's global extension-tracing gate to the current tracing + // state. Covers static ZE_ENABLE_TRACING_LAYER enablement (the driver becomes + // usable only by the time the app registers) and drivers registered on after + // a dynamic enable. Disabling is handled centrally by zelDisableTracingLayer + // (respecting sticky env), so only propagate the enable here. + bool tracingOn = ze_lib::context->tracingLayerEnableCounter.load() > 0; + if( !tracingOn && ze_lib::context->loaderContext ) + tracingOn = ze_lib::context->loaderContext->tracingLayerEnabled; + if( tracingOn ) { + void* pfnEnableRaw = nullptr; + if( ZE_RESULT_SUCCESS == pfnGetExtensionFunctionAddress( + hDriver, "zelDriverEnableTracing", &pfnEnableRaw ) && pfnEnableRaw ) { + reinterpret_cast(pfnEnableRaw)( hDriver, true ); + } + } + + auto pfnSet = reinterpret_cast( pfnRaw ); + return pfnSet( hDriver, functionName, pUserData, prologue, epilogue ); +} + } //extern "c" \ No newline at end of file diff --git a/source/lib/ze_lib.h b/source/lib/ze_lib.h index ce457d59..8c23167f 100644 --- a/source/lib/ze_lib.h +++ b/source/lib/ze_lib.h @@ -28,6 +28,10 @@ #include #include +// Forward declaration: the full definition lives in ze_loader_internal.h, which +// ze_lib.cpp includes. Here we only need it for a pointer member. +namespace loader { class context_t; } + namespace ze_lib { /////////////////////////////////////////////////////////////////////////////// @@ -184,6 +188,10 @@ namespace ze_lib bool debugTraceEnabled = false; bool dynamicTracingSupported = true; ze_pfnDriverGet_t loaderDriverGet = nullptr; + // Loader context, resolved in a build-portable way during Init: directly + // in the dynamic build, or via zelLoaderGetContext() in the static build + // (where the loader symbols are not link-time visible). May be null. + loader::context_t *loaderContext = nullptr; std::atomic teardownCallbacksCount{0}; std::map teardownCallbacks; std::mutex teardownCallbacksMutex; diff --git a/source/loader/ze_loader.cpp b/source/loader/ze_loader.cpp index 14a0ef27..725c3108 100644 --- a/source/loader/ze_loader.cpp +++ b/source/loader/ze_loader.cpp @@ -8,6 +8,7 @@ #include "ze_loader_utils.h" #include "driver_discovery.h" +#include "../lib/ze_lib.h" #include #include @@ -478,6 +479,20 @@ namespace loader return true; } + ze_result_t enableDriverExtensionTracing(driver_t &driver, ze_bool_t enable) { + auto pfnGetExtensionFunctionAddress = driver.dditable.ze.Driver.pfnGetExtensionFunctionAddress; + if (nullptr == pfnGetExtensionFunctionAddress) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + + void *pfnRaw = nullptr; + // Global driver-level hook; the handle is not needed to resolve it. + ze_result_t res = pfnGetExtensionFunctionAddress(nullptr, "zelDriverEnableTracing", &pfnRaw); + if (res != ZE_RESULT_SUCCESS || nullptr == pfnRaw) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; // driver doesn't support it; skip + + return reinterpret_cast(pfnRaw)(nullptr, enable); + } + ze_result_t context_t::init_driver(driver_t &driver, ze_init_flags_t flags, ze_init_driver_type_desc_t* desc) { bool loadDriver = false; if (debugTraceEnabled) { @@ -596,6 +611,17 @@ namespace loader return ZE_RESULT_ERROR_UNINITIALIZED; } + // If the tracing layer is already enabled (statically via + // ZE_ENABLE_TRACING_LAYER, or dynamically via zelEnableTracingLayer), + // propagate the enable to this now-usable driver so env-enabled and + // late-loaded drivers participate in extension-function tracing. Done + // here (rather than inside the DDI-init block above) because some drivers + // populate their dispatch tables during discovery and skip that block. + if (tracingLayerEnabled || + (ze_lib::context && ze_lib::context->tracingLayerEnableCounter.load() > 0)) { + enableDriverExtensionTracing(driver, true); + } + return ZE_RESULT_SUCCESS; } diff --git a/source/loader/ze_loader_internal.h b/source/loader/ze_loader_internal.h index 45dce29c..bf0cf1f3 100644 --- a/source/loader/ze_loader_internal.h +++ b/source/loader/ze_loader_internal.h @@ -184,4 +184,10 @@ namespace loader extern ze_handle_t* loaderDispatch; extern zer_dditable_t* defaultZerDdiTable; extern context_t *context; + + // Enable/disable extension-function tracing on a single driver by resolving + // its "zelDriverEnableTracing" hook by name. No-op (returns UNSUPPORTED) for + // drivers that don't implement it. Used to propagate the tracing-layer + // enable/disable state (env + dynamic) down to each driver. + ze_result_t enableDriverExtensionTracing(driver_t &driver, ze_bool_t enable); } diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index a1fa8dd6..45394fd2 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -7,6 +7,7 @@ add_executable( loader_validation_layer.cpp driver_ordering_helper_tests.cpp loader_tracing_layer.cpp + loader_ext_fn_callback.cpp ) # Only include driver_ordering_unit_tests and driver_teardown_unit_tests for static builds or non-Windows platforms @@ -264,6 +265,15 @@ set_property(TEST tests_tracing_layer_state_enabled_via_environment_and_dynamic add_test(NAME tests_tracing_layer_state_enabled_via_environment_disable_dynamic COMMAND tests --gtest_filter=*TracingLayerState.GivenTracingLayerEnabledViaEnvironmentAndDynamicallyWhenDisablingDynamicTracingThenStateRemainsTrue) set_property(TEST tests_tracing_layer_state_enabled_via_environment_disable_dynamic PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1") +# Extension-function callback (zelDriverSetExtensionFunctionCallback) tests +# Dynamic control suite (toggles tracing at runtime). +add_test(NAME tests_ext_fn_callback COMMAND tests --gtest_filter=*ExtFnCallback.*) +set_property(TEST tests_ext_fn_callback PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1") + +# Static-enablement suite (ZE_ENABLE_TRACING_LAYER=1, sticky for the process). +add_test(NAME tests_ext_fn_callback_env COMMAND tests --gtest_filter=*ExtFnCallbackEnviron.*) +set_property(TEST tests_ext_fn_callback_env PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1;ZE_ENABLE_TRACING_LAYER=1") + add_test(NAME test_zello_world_legacy COMMAND zello_world --enable_legacy_init --enable_null_driver --force_loader_intercepts --enable_validation_layer --enable_tracing_layer --enable_tracing_layer_runtime) set_property(TEST test_zello_world_legacy PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1") @@ -494,8 +504,18 @@ foreach(test_name IN ITEMS LoaderTranslateHandles.GivenLevelZeroLoaderPresentWhenCallingZelLoaderTranslateHandleInternalWithInterceptEnabledAndDDiSupportEnabledThenExpectNoHandleTranslationForDriver LoaderTranslateHandles.GivenLevelZeroLoaderPresentWhenCallingZelLoaderTranslateHandleInternalWithInterceptEnabledAndDDiSupportDisabledThenExpectHandleTranslationForDevice LoaderTranslateHandles.GivenLevelZeroLoaderPresentWhenCallingZelLoaderTranslateHandleInternalWithInterceptEnabledAndDDiSupportEnabledThenExpectNoHandleTranslationForDevice) - add_test(NAME ${test_name}_alt_drivers COMMAND tests --gtest_filter=*${test_name}) - set_property(TEST ${test_name}_alt_drivers APPEND PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1;${ALT_DRIVERS_ENV}") + # Derive a short, convention-matching ctest name from the long gtest name + # (e.g. ...ThenExpectHandleTranslationForCommandQueue -> command_queue). + string(REGEX REPLACE ".*ThenExpect(No)?HandleTranslationFor" "" _handle_type "${test_name}") + string(REGEX REPLACE "([a-z0-9])([A-Z])" "\\1_\\2" _handle_type "${_handle_type}") + string(TOLOWER "${_handle_type}" _handle_type) + if(test_name MATCHES "NoHandleTranslation") + set(_short_name "tests_loader_translate_handles_${_handle_type}_ddi_enabled_alt_drivers") + else() + set(_short_name "tests_loader_translate_handles_${_handle_type}_alt_drivers") + endif() + add_test(NAME ${_short_name} COMMAND tests --gtest_filter=*${test_name}) + set_property(TEST ${_short_name} APPEND PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1;${ALT_DRIVERS_ENV}") endforeach() add_test(NAME tests_single_driver_sysman_vf_management_api COMMAND tests --gtest_filter=*GivenLevelZeroLoaderPresentWhenCallingSysManVfApisThenExpectNullDriverIsReachedSuccessfully) diff --git a/test/loader_ext_fn_callback.cpp b/test/loader_ext_fn_callback.cpp new file mode 100644 index 00000000..001cdf37 --- /dev/null +++ b/test/loader_ext_fn_callback.cpp @@ -0,0 +1,316 @@ +/* + * + * Copyright (C) 2026 Intel Corporation + * + * SPDX-License-Identifier: MIT + * + */ + +#include "gtest/gtest.h" + +#include "loader/ze_loader.h" +#include "ze_api.h" + +#include + +namespace { + +// Signature of the sample extension function exposed by the null driver. +typedef ze_result_t (ZE_APICALL *pfnSampleExtFunc_t)( + ze_driver_handle_t, uint32_t, uint32_t*); + +constexpr uintptr_t kInstanceSentinel = 0xABCD1234u; + +// State the prologue/epilogue callbacks record into, reached via pTracerUserData. +struct CallbackState { + int prologCount = 0; + int epilogCount = 0; + void* prologUserData = nullptr; + void* epilogUserData = nullptr; + ze_result_t epilogResult = ZE_RESULT_FORCE_UINT32; + uintptr_t instanceValueSeenInEpilog = 0; + bool prologRanBeforeEpilog = false; +}; + +void ZE_APICALL prologueCb(void* /*pParams*/, ze_result_t /*result*/, + void* pTracerUserData, void** ppTracerInstanceUserData) { + auto* s = static_cast(pTracerUserData); + s->prologCount++; + s->prologUserData = pTracerUserData; + *ppTracerInstanceUserData = reinterpret_cast(kInstanceSentinel); +} + +void ZE_APICALL epilogueCb(void* /*pParams*/, ze_result_t result, + void* pTracerUserData, void** ppTracerInstanceUserData) { + auto* s = static_cast(pTracerUserData); + s->epilogCount++; + s->epilogUserData = pTracerUserData; + s->epilogResult = result; + s->prologRanBeforeEpilog = (s->prologCount == 1); + s->instanceValueSeenInEpilog = + reinterpret_cast(*ppTracerInstanceUserData); +} + +ze_driver_handle_t getFirstDriver() { + EXPECT_EQ(ZE_RESULT_SUCCESS, zeInit(0)); + uint32_t count = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, zeDriverGet(&count, nullptr)); + EXPECT_GT(count, 0u); + count = 1; + ze_driver_handle_t hDriver = nullptr; + EXPECT_EQ(ZE_RESULT_SUCCESS, zeDriverGet(&count, &hDriver)); + EXPECT_NE(nullptr, hDriver); + return hDriver; +} + +pfnSampleExtFunc_t getSampleExtFunc(ze_driver_handle_t hDriver) { + void* addr = nullptr; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zeDriverGetExtensionFunctionAddress(hDriver, "zeSampleExtFunc", &addr)); + EXPECT_NE(nullptr, addr); + return reinterpret_cast(addr); +} + +void unregister(ze_driver_handle_t hDriver, const char* name) { + zelDriverSetExtensionFunctionCallback(hDriver, name, nullptr, nullptr, nullptr); +} + +// --------------------------------------------------------------------------- +// Dynamic control suite: tracing is toggled at runtime via +// zelEnableTracingLayer/zelDisableTracingLayer (no ZE_ENABLE_TRACING_LAYER env). +// Each test balances enable/disable and unregisters so process state stays clean. +// --------------------------------------------------------------------------- + +TEST(ExtFnCallback, PrologueAndEpilogueFireOnCall) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 21, &out)); + + EXPECT_EQ(42u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(1, state.epilogCount); + EXPECT_EQ(&state, state.prologUserData); + EXPECT_EQ(&state, state.epilogUserData); + EXPECT_TRUE(state.prologRanBeforeEpilog); + EXPECT_EQ(ZE_RESULT_SUCCESS, state.epilogResult); + EXPECT_EQ(kInstanceSentinel, state.instanceValueSeenInEpilog); + + unregister(hDriver, "zeSampleExtFunc"); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, RegisterBeforeFetchStillFires) { + auto hDriver = getFirstDriver(); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 5, &out)); + + EXPECT_EQ(10u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(1, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, UnregisterStopsCallbacks) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, nullptr, nullptr)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 7, &out)); + + EXPECT_EQ(14u, out); + EXPECT_EQ(0, state.prologCount); + EXPECT_EQ(0, state.epilogCount); + + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, UnknownFunctionNameNeverFires) { + auto hDriver = getFirstDriver(); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback( + hDriver, "zeNeverImplementedExtFunc", &state, prologueCb, epilogueCb)); + + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 3, &out)); + + EXPECT_EQ(6u, out); + EXPECT_EQ(0, state.prologCount); + EXPECT_EQ(0, state.epilogCount); + + unregister(hDriver, "zeNeverImplementedExtFunc"); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, NullArgumentsReturnErrors) { + auto hDriver = getFirstDriver(); + + EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_HANDLE, + zelDriverSetExtensionFunctionCallback(nullptr, "zeSampleExtFunc", + nullptr, prologueCb, epilogueCb)); + EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_POINTER, + zelDriverSetExtensionFunctionCallback(hDriver, nullptr, nullptr, + prologueCb, epilogueCb)); +} + +// Two-level gate: registered but tracing layer NOT enabled -> must not fire. +TEST(ExtFnCallback, NotEnabledDoesNotFire) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 9, &out)); + + EXPECT_EQ(18u, out); // body runs + EXPECT_EQ(0, state.prologCount); // gate closed -> no callbacks + EXPECT_EQ(0, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); +} + +// Disabling the tracing layer stops callbacks even while still registered. +TEST(ExtFnCallback, DisableStopsCallbacks) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 1, &out)); + EXPECT_EQ(1, state.prologCount); // fires while enabled + + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 1, &out)); + EXPECT_EQ(1, state.prologCount); // no additional fire after disable + EXPECT_EQ(1, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); +} + +// Underflow guard (fix #1): a disable with no matching enable must be a safe +// no-op. If the unsigned counter had underflowed, the subsequent enable would +// not detect the 0->1 edge, the driver would never be enabled, and the callback +// would not fire - so a passing "fires" assertion proves no corruption occurred. +TEST(ExtFnCallback, DisableWithoutEnableIsSafe) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + // Counter is 0 here (all prior tests balanced). Extra disables must no-op. + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 2, &out)); + EXPECT_EQ(1, state.prologCount); // enable's 0->1 edge still worked + EXPECT_EQ(1, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// --------------------------------------------------------------------------- +// Environment suite: run with ZE_ENABLE_TRACING_LAYER=1 (separate ctest entry). +// Tracing is enabled statically at init; the app never calls +// zelEnableTracingLayer, and per documented behavior it stays enabled for the +// whole process. +// --------------------------------------------------------------------------- + +TEST(ExtFnCallbackEnviron, EnvKeepsTracingEnabled) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + // No zelEnableTracingLayer call: the driver was enabled at init (Site A). + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 4, &out)); + + EXPECT_EQ(8u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(1, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); +} + +// A spurious disable under static enablement must not turn tracing off (sticky +// env) and must not corrupt state (underflow guard). +TEST(ExtFnCallbackEnviron, EnvDisableIsNoOp) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); // counter 0 -> no-op + + CallbackState state; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", + &state, prologueCb, epilogueCb)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 6, &out)); + + EXPECT_EQ(12u, out); + EXPECT_EQ(1, state.prologCount); // still fires: env-enabled tracing is sticky + EXPECT_EQ(1, state.epilogCount); + + unregister(hDriver, "zeSampleExtFunc"); +} + +} // namespace From b3bd4bc032b3225e019f22e0528c6fb0f3832e9c Mon Sep 17 00:00:00 2001 From: Russell McGuire Date: Tue, 21 Jul 2026 22:11:57 -0700 Subject: [PATCH 2/2] 2nd version of driver extension tracing + Convert to include use of Tracer handles Signed-off-by: Russell McGuire --- include/loader/ze_loader.h | 71 ++++++--- source/drivers/null/ze_null.cpp | 42 ++--- source/drivers/null/ze_null.h | 37 +++-- source/layers/tracing/tracing.h | 5 + source/layers/tracing/tracing_imp.cpp | 144 +++++++++++++++++ source/layers/tracing/tracing_imp.h | 56 +++++++ source/layers/tracing/ze_tracing.cpp | 56 +++++++ source/lib/ze_lib.cpp | 74 ++------- source/lib/zel_tracing_libapi.cpp | 31 ++++ test/CMakeLists.txt | 2 +- test/loader_ext_fn_callback.cpp | 214 +++++++++++++++++++------- 11 files changed, 555 insertions(+), 177 deletions(-) diff --git a/include/loader/ze_loader.h b/include/loader/ze_loader.h index cfc89372..8b8cf487 100644 --- a/include/loader/ze_loader.h +++ b/include/loader/ze_loader.h @@ -13,6 +13,7 @@ #endif #include "../ze_api.h" +#include "../layers/zel_tracing_register_cb.h" #if !defined(__cplusplus) #include @@ -602,39 +603,67 @@ typedef ze_result_t (ZE_APICALL *zel_pfnDriverEnableTracing_t)( ); /////////////////////////////////////////////////////////////////////////////// -/// @brief Registers prologue/epilogue callbacks for a named extension function. +/// @brief Signature of the per-driver hook the loader/tracing-layer uses to +/// install its extension-function interception wrappers on a driver. +/// +/// A driver that supports extension-function tracing exposes this by name +/// ("zelDriverSetLoaderCallbackForExtension") via +/// zeDriverGetExtensionFunctionAddress. The tracing layer calls it to register a +/// single loader-owned prologue/epilogue wrapper (plus an opaque loader context) +/// for the named extension function. The driver invokes @p loaderPrologue before, +/// and @p loaderEpilogue after, the body of the extension function named +/// @p functionName, forwarding @p pLoaderContext back unchanged. Passing null for +/// both wrappers unregisters. The loader owns the fan-out to any number of +/// registered tracers, so the driver stores at most one wrapper per function. +typedef ze_result_t (ZE_APICALL *zel_pfnDriverSetLoaderCallbackForExtension_t)( + ze_driver_handle_t hDriver, // [in] handle of the driver instance + const char* functionName, // [in] extension function name to intercept + zel_pfnDriverExtensionFunctionCb_t loaderPrologue, // [in][optional] loader prologue wrapper + zel_pfnDriverExtensionFunctionCb_t loaderEpilogue, // [in][optional] loader epilogue wrapper + void* pLoaderContext // [in][optional] loader context echoed to wrappers + ); + +/////////////////////////////////////////////////////////////////////////////// +/// @brief Registers a prologue or epilogue callback on a tracer for a named +/// extension function of a specific driver. /// /// Extension functions obtained by string name via /// zeDriverGetExtensionFunctionAddress() return a raw driver pointer that the -/// application calls directly, bypassing the loader and therefore the tracing -/// layer. This API provides a driver-side interception hook: the specified -/// driver invokes the registered @p prologue before, and @p epilogue after, the -/// body of the extension function named @p functionName. +/// application calls directly, bypassing the loader and therefore the per-API +/// tracing interceptors. This API routes such functions through the same tracer +/// (::zel_tracer_handle_t) infrastructure used for core APIs: the tracing layer +/// installs a loader-owned wrapper on @p hDriver (via the driver's +/// zelDriverSetLoaderCallbackForExtension hook) and fans out to every enabled +/// tracer that registered @p functionName for @p hDriver. /// -/// Registration is keyed by @p functionName and is order-independent relative to -/// zeDriverGetExtensionFunctionAddress() — it takes effect on the next invocation -/// of the function, even if the application already cached the function pointer. -/// Passing null for both @p prologue and @p epilogue unregisters the callbacks. +/// Registration is keyed by (@p hDriver, @p functionName) and is order-independent +/// relative to zeDriverGetExtensionFunctionAddress() — it takes effect on the next +/// invocation even if the application already cached the function pointer. The +/// callback receives the tracer's pUserData (from ::zelTracerCreate) as +/// pTracerUserData. Multiple tracers may register the same function to stack +/// callbacks. The callbacks fire only when the tracing layer is enabled for the +/// driver and the tracer is enabled. /// -/// @param[in] hDriver handle of the driver instance +/// @param[in] hTracer handle of the tracer to register the callback on +/// @param[in] hDriver handle of the driver whose extension function to trace /// @param[in] functionName name of the extension function to intercept -/// @param[in] pUserData user data passed to the callbacks as pTracerUserData -/// @param[in] prologue handler invoked before the function body (may be null) -/// @param[in] epilogue handler invoked after the function body (may be null) +/// @param[in] callback_type ::ZEL_REGISTER_PROLOGUE or ::ZEL_REGISTER_EPILOGUE +/// @param[in] pCallback handler to register (null clears that slot) /// /// @return -/// - ZE_RESULT_SUCCESS on success (including unregister). -/// - ZE_RESULT_ERROR_UNINITIALIZED if the loader is not initialized. -/// - ZE_RESULT_ERROR_UNSUPPORTED_FEATURE if the driver does not implement this hook. -/// - ZE_RESULT_ERROR_INVALID_NULL_HANDLE if @p hDriver is null. +/// - ZE_RESULT_SUCCESS on success (including clearing a slot). +/// - ZE_RESULT_ERROR_UNINITIALIZED if the loader/tracing layer is not initialized. +/// - ZE_RESULT_ERROR_UNSUPPORTED_FEATURE if the driver does not implement the hook. +/// - ZE_RESULT_ERROR_INVALID_NULL_HANDLE if @p hTracer or @p hDriver is null. /// - ZE_RESULT_ERROR_INVALID_NULL_POINTER if @p functionName is null. +/// - ZE_RESULT_ERROR_INVALID_ARGUMENT if the tracer is not in the disabled state. ZE_DLLEXPORT ze_result_t ZE_APICALL -zelDriverSetExtensionFunctionCallback( +zelTracerDriverExtensionRegisterCallback( + zel_tracer_handle_t hTracer, // [in] handle of the tracer ze_driver_handle_t hDriver, // [in] handle of the driver instance const char* functionName, // [in] extension function name to intercept - void* pUserData, // [in][optional] user data passed to callbacks - zel_pfnDriverExtensionFunctionCb_t prologue, // [in][optional] prologue handler - zel_pfnDriverExtensionFunctionCb_t epilogue // [in][optional] epilogue handler + zel_tracer_reg_t callback_type, // [in] prologue or epilogue + zel_pfnDriverExtensionFunctionCb_t pCallback // [in][optional] handler (null clears slot) ); #if defined(__cplusplus) diff --git a/source/drivers/null/ze_null.cpp b/source/drivers/null/ze_null.cpp index 832e8a86..ae2e45f3 100644 --- a/source/drivers/null/ze_null.cpp +++ b/source/drivers/null/ze_null.cpp @@ -58,8 +58,8 @@ namespace driver { if( nullptr == name || nullptr == ppFunctionAddress ) return ZE_RESULT_ERROR_INVALID_NULL_POINTER; - if( 0 == strcmp( name, "zelDriverSetExtensionFunctionCallback" ) ) { - *ppFunctionAddress = reinterpret_cast( &driver::zelDriverSetExtensionFunctionCallback ); + if( 0 == strcmp( name, "zelDriverSetLoaderCallbackForExtension" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zelDriverSetLoaderCallbackForExtension ); return ZE_RESULT_SUCCESS; } if( 0 == strcmp( name, "zelDriverEnableTracing" ) ) { @@ -713,8 +713,8 @@ namespace driver ze_result_t ZE_APICALL zeSampleExtFunc( ze_driver_handle_t hDriver, uint32_t input, uint32_t* pOutput ) { - // Snapshot any registered callbacks for this function (name-keyed). - context_t::extension_function_callbacks_t cbs; + // Snapshot the single loader wrapper registered for this function. + context_t::loader_extension_callbacks_t cbs; bool haveCbs = false; { std::lock_guard lock( context.extensionCallbackMutex ); @@ -725,8 +725,8 @@ namespace driver } } - // Two-level gate: callbacks fire only when tracing is globally enabled - // AND a callback is registered for this function. + // Two-level gate: the wrapper fires only when tracing is globally enabled + // AND a loader wrapper is registered for this function. const bool fire = haveCbs && context.extensionCallbacksEnabled.load(); // Typed parameter block the driver exposes to the callbacks. @@ -734,15 +734,15 @@ namespace driver void* pInstanceData = nullptr; ze_result_t result = ZE_RESULT_SUCCESS; - if( fire && nullptr != cbs.prologue ) - cbs.prologue( ¶ms, result, cbs.pUserData, &pInstanceData ); + if( fire && nullptr != cbs.loaderPrologue ) + cbs.loaderPrologue( ¶ms, result, cbs.pLoaderContext, &pInstanceData ); // The (trivial) work of the extension function. if( nullptr != pOutput ) *pOutput = input * 2; - if( fire && nullptr != cbs.epilogue ) - cbs.epilogue( ¶ms, result, cbs.pUserData, &pInstanceData ); + if( fire && nullptr != cbs.loaderEpilogue ) + cbs.loaderEpilogue( ¶ms, result, cbs.pLoaderContext, &pInstanceData ); return result; } @@ -759,24 +759,26 @@ namespace driver } /////////////////////////////////////////////////////////////////////////// - /// @brief Driver-side registration entry (resolved by name from the loader). - /// Permissive and name-keyed: any name registers; null+null unregisters. - ze_result_t ZE_APICALL zelDriverSetExtensionFunctionCallback( - ze_driver_handle_t, const char* functionName, void* pUserData, - zel_pfnDriverExtensionFunctionCb_t prologue, - zel_pfnDriverExtensionFunctionCb_t epilogue ) + /// @brief Driver-side loader-callback registration entry (resolved by name + /// from the tracing layer). Stores the single loader wrapper (+ opaque + /// context) per function name; null+null unregisters. + ze_result_t ZE_APICALL zelDriverSetLoaderCallbackForExtension( + ze_driver_handle_t, const char* functionName, + zel_pfnDriverExtensionFunctionCb_t loaderPrologue, + zel_pfnDriverExtensionFunctionCb_t loaderEpilogue, + void* pLoaderContext ) { if( nullptr == functionName ) return ZE_RESULT_ERROR_INVALID_NULL_POINTER; std::lock_guard lock( context.extensionCallbackMutex ); - if( nullptr == prologue && nullptr == epilogue ) { + if( nullptr == loaderPrologue && nullptr == loaderEpilogue ) { context.extensionCallbacks.erase( functionName ); } else { auto& entry = context.extensionCallbacks[ functionName ]; - entry.pUserData = pUserData; - entry.prologue = prologue; - entry.epilogue = epilogue; + entry.loaderPrologue = loaderPrologue; + entry.loaderEpilogue = loaderEpilogue; + entry.pLoaderContext = pLoaderContext; } return ZE_RESULT_SUCCESS; } diff --git a/source/drivers/null/ze_null.h b/source/drivers/null/ze_null.h index f0849b07..99161f29 100644 --- a/source/drivers/null/ze_null.h +++ b/source/drivers/null/ze_null.h @@ -53,16 +53,19 @@ namespace driver bool ddiExtensionSupported = false; std::vector env_vars{}; - // Registry for zelDriverSetExtensionFunctionCallback: maps an extension - // function name to the prologue/epilogue the driver invokes from inside - // that function's body. Keyed by name (order-independent vs fetch). - struct extension_function_callbacks_t { - void* pUserData = nullptr; - zel_pfnDriverExtensionFunctionCb_t prologue = nullptr; - zel_pfnDriverExtensionFunctionCb_t epilogue = nullptr; + // Registry for zelDriverSetLoaderCallbackForExtension: maps an extension + // function name to the single loader-owned wrapper the driver invokes + // from that function's body. The loader/tracing-layer owns the fan-out to + // any number of tracers, so the driver stores at most one wrapper (plus + // an opaque loader context) per function. Keyed by name (order- + // independent vs fetch). + struct loader_extension_callbacks_t { + zel_pfnDriverExtensionFunctionCb_t loaderPrologue = nullptr; + zel_pfnDriverExtensionFunctionCb_t loaderEpilogue = nullptr; + void* pLoaderContext = nullptr; }; std::mutex extensionCallbackMutex; - std::map extensionCallbacks; + std::map extensionCallbacks; // Global gate for extension-function callbacks, toggled by the loader via // zelDriverEnableTracing. Callbacks fire only when this is set AND a @@ -95,8 +98,8 @@ namespace driver // Extension-function callback prototype demonstration. // // "zeSampleExtFunc" is a stand-in vendor extension function reachable only by - // name via zeDriverGetExtensionFunctionAddress. Its body invokes any - // prologue/epilogue registered through zelDriverSetExtensionFunctionCallback, + // name via zeDriverGetExtensionFunctionAddress. Its body invokes the single + // loader-owned wrapper registered through zelDriverSetLoaderCallbackForExtension, // passing a typed params block (the driver knows its own signature). typedef struct _ze_sample_ext_func_params_t { @@ -108,12 +111,14 @@ namespace driver ze_result_t ZE_APICALL zeSampleExtFunc( ze_driver_handle_t hDriver, uint32_t input, uint32_t* pOutput ); - // Driver-side registration entry, resolved by name from the loader's - // zelDriverSetExtensionFunctionCallback forward. - ze_result_t ZE_APICALL zelDriverSetExtensionFunctionCallback( - ze_driver_handle_t hDriver, const char* functionName, void* pUserData, - zel_pfnDriverExtensionFunctionCb_t prologue, - zel_pfnDriverExtensionFunctionCb_t epilogue ); + // Driver-side loader-callback registration entry, resolved by name from the + // tracing layer. Stores the single loader wrapper (+ context) per function; + // null+null unregisters. + ze_result_t ZE_APICALL zelDriverSetLoaderCallbackForExtension( + ze_driver_handle_t hDriver, const char* functionName, + zel_pfnDriverExtensionFunctionCb_t loaderPrologue, + zel_pfnDriverExtensionFunctionCb_t loaderEpilogue, + void* pLoaderContext ); // Driver-side enable/disable of extension-function callbacks, resolved by // name from the loader when the tracing layer is enabled/disabled. diff --git a/source/layers/tracing/tracing.h b/source/layers/tracing/tracing.h index d6c296c8..32e0acc0 100644 --- a/source/layers/tracing/tracing.h +++ b/source/layers/tracing/tracing.h @@ -10,6 +10,7 @@ #include "ze_api.h" #include "layers/zel_tracing_api.h" #include "layers/zel_tracing_register_cb.h" +#include "loader/ze_loader.h" #include "ze_tracing_cb_structs.h" #include "zer_tracing_cb_structs.h" @@ -30,6 +31,10 @@ struct APITracer : _zel_tracer_handle_t { virtual zel_zer_all_callbacks_t& getZerProEpilogues(zel_tracer_reg_t callback_type, ze_result_t& result) = 0; virtual ze_result_t resetAllCallbacks() = 0; virtual ze_result_t enableTracer(ze_bool_t enable) = 0; + virtual ze_result_t registerExtensionCallback(ze_driver_handle_t hDriver, + const char *functionName, + zel_tracer_reg_t callback_type, + zel_pfnDriverExtensionFunctionCb_t pCallback) = 0; }; ze_result_t createAPITracer(const zel_tracer_desc_t *desc, zel_tracer_handle_t *phTracer); diff --git a/source/layers/tracing/tracing_imp.cpp b/source/layers/tracing/tracing_imp.cpp index 35bd12bb..43f0f01f 100644 --- a/source/layers/tracing/tracing_imp.cpp +++ b/source/layers/tracing/tracing_imp.cpp @@ -118,6 +118,7 @@ ze_result_t APITracerImp::resetAllCallbacks() { this->tracerFunctions.coreEpilogues = {}; this->tracerFunctions.runtimePrologues = {}; this->tracerFunctions.runtimeEpilogues = {}; + this->tracerFunctions.extensionCallbacks.clear(); return ZE_RESULT_SUCCESS; } @@ -126,6 +127,30 @@ ze_result_t APITracerImp::enableTracer(ze_bool_t enable) { return pGlobalAPITracerContextImp->enableTracingImp(this, enable); } +ze_result_t APITracerImp::registerExtensionCallback( + ze_driver_handle_t hDriver, const char *functionName, + zel_tracer_reg_t callback_type, + zel_pfnDriverExtensionFunctionCb_t pCallback) { + + // Mirror the per-API register callbacks: registration is only permitted while + // the tracer is disabled so the active tracer array is never mutated live. + if (this->tracingState != disabledState) + return ZE_RESULT_ERROR_INVALID_ARGUMENT; + + ExtensionFunctionKey key{hDriver, functionName}; + auto &entry = this->tracerFunctions.extensionCallbacks[key]; + if (callback_type == ZEL_REGISTER_PROLOGUE) + entry.prologue = pCallback; + else + entry.epilogue = pCallback; + + // Drop fully-cleared entries so the fan-out never carries dead keys. + if (entry.prologue == nullptr && entry.epilogue == nullptr) + this->tracerFunctions.extensionCallbacks.erase(key); + + return ZE_RESULT_SUCCESS; +} + void APITracerImp::copyCoreCbsToAllCbs(zel_ze_all_callbacks_t& allCbs, zel_core_callbacks_t& cbs) { allCbs.Global.pfnInitCb = cbs.Global.pfnInitCb; @@ -497,4 +522,123 @@ void APITracerContextImp::releaseActivetracersList() { nullptr, std::memory_order_relaxed); } +namespace { +// Process-lifetime registry of loader contexts, one per (hDriver, functionName). +// Entries are never erased so the pointer handed to a driver stays valid for the +// life of the process; std::map guarantees stable node addresses across inserts. +std::mutex loaderExtensionContextMutex; +std::map loaderExtensionContexts; + +// Per-call scratch carried from the loader prologue wrapper to the epilogue +// wrapper via the driver's ppTracerInstanceUserData slot. Holds a snapshot of +// each participating tracer's epilogue callback plus its per-call instance data, +// so the epilogue never dereferences the (possibly retired) active tracer array. +struct ExtensionCallFrame { + std::vector> + epilogCallbacks; + std::vector instanceUserData; +}; +} // namespace + +LoaderExtensionContext * +getOrCreateLoaderExtensionContext(ze_driver_handle_t hDriver, + const char *functionName) { + ExtensionFunctionKey key{hDriver, functionName}; + std::lock_guard lock(loaderExtensionContextMutex); + auto it = loaderExtensionContexts.find(key); + if (it == loaderExtensionContexts.end()) { + it = loaderExtensionContexts + .emplace(key, LoaderExtensionContext{hDriver, functionName}) + .first; + } + return &it->second; +} + +void ZE_APICALL loaderExtensionPrologue(void *pParams, ze_result_t result, + void *pLoaderContext, + void **ppTracerInstanceUserData) { + if (ppTracerInstanceUserData != nullptr) + *ppTracerInstanceUserData = nullptr; + if (pLoaderContext == nullptr || ppTracerInstanceUserData == nullptr) + return; + + // Recursion guard: suppress nested extension tracing on this thread. The flag + // is owned per-phase (set here, cleared before the driver body runs) so the + // body's own core-API calls remain traceable; the epilogue re-establishes it. + if (tracingInProgress) + return; + tracingInProgress = 1; + + auto *ctx = static_cast(pLoaderContext); + ExtensionFunctionKey key{ctx->hDriver, ctx->functionName}; + + std::vector> + prologCallbacks; + auto *frame = new ExtensionCallFrame(); + + tracer_array_t *currentTracerArray = + (tracer_array_t *)pGlobalAPITracerContextImp->getActiveTracersList(); + if (currentTracerArray && currentTracerArray->tracerArrayCount) { + for (size_t i = 0; i < currentTracerArray->tracerArrayCount; i++) { + auto &tracerEntry = currentTracerArray->tracerArrayEntries[i]; + auto cbIt = tracerEntry.extensionCallbacks.find(key); + if (cbIt == tracerEntry.extensionCallbacks.end()) + continue; + // Push prologue and epilogue together so indices stay aligned per + // tracer and share one instance-data slot across both phases. + APITracerCallbackStateImp prolog; + prolog.current_api_callback = cbIt->second.prologue; + prolog.pUserData = tracerEntry.pUserData; + prologCallbacks.push_back(prolog); + APITracerCallbackStateImp epilog; + epilog.current_api_callback = cbIt->second.epilogue; + epilog.pUserData = tracerEntry.pUserData; + frame->epilogCallbacks.push_back(epilog); + } + } + pGlobalAPITracerContextImp->releaseActivetracersList(); + + // No participating tracer: free the frame and leave a null instance handle so + // the epilogue wrapper is a no-op. + if (prologCallbacks.empty()) { + delete frame; + tracingInProgress = 0; + return; + } + + frame->instanceUserData.resize(prologCallbacks.size(), nullptr); + for (size_t i = 0; i < prologCallbacks.size(); i++) { + if (prologCallbacks[i].current_api_callback != nullptr) + prologCallbacks[i].current_api_callback( + pParams, result, prologCallbacks[i].pUserData, + &frame->instanceUserData[i]); + } + + *ppTracerInstanceUserData = frame; + tracingInProgress = 0; +} + +void ZE_APICALL loaderExtensionEpilogue(void *pParams, ze_result_t result, + void * /*pLoaderContext*/, + void **ppTracerInstanceUserData) { + if (ppTracerInstanceUserData == nullptr) + return; + auto *frame = + static_cast(*ppTracerInstanceUserData); + if (frame == nullptr) + return; + + tracingInProgress = 1; + for (size_t i = 0; i < frame->epilogCallbacks.size(); i++) { + if (frame->epilogCallbacks[i].current_api_callback != nullptr) + frame->epilogCallbacks[i].current_api_callback( + pParams, result, frame->epilogCallbacks[i].pUserData, + &frame->instanceUserData[i]); + } + tracingInProgress = 0; + + delete frame; + *ppTracerInstanceUserData = nullptr; +} + } // namespace tracing_layer diff --git a/source/layers/tracing/tracing_imp.h b/source/layers/tracing/tracing_imp.h index c03be8be..0e8ad2f7 100644 --- a/source/layers/tracing/tracing_imp.h +++ b/source/layers/tracing/tracing_imp.h @@ -10,12 +10,15 @@ #include "tracing.h" #include "ze_api.h" #include "ze_tracing_cb_structs.h" +#include "loader/ze_loader.h" #include #include #include #include +#include #include +#include #include #include @@ -32,12 +35,42 @@ namespace tracing_layer { extern thread_local ze_bool_t tracingInProgress; extern struct APITracerContextImp *pGlobalAPITracerContextImp; +// Identifies a traced extension function on a specific driver. Registration and +// per-call fan-out are keyed by (hDriver, functionName) so a callback registered +// for one driver never fires for another driver's same-named function. +struct ExtensionFunctionKey { + ze_driver_handle_t hDriver; + std::string functionName; + bool operator<(const ExtensionFunctionKey &rhs) const { + if (hDriver != rhs.hDriver) + return hDriver < rhs.hDriver; + return functionName < rhs.functionName; + } +}; + +// A single tracer's prologue/epilogue for one extension function. +struct ExtensionFunctionCallbacks { + zel_pfnDriverExtensionFunctionCb_t prologue = nullptr; + zel_pfnDriverExtensionFunctionCb_t epilogue = nullptr; +}; + +// Loader-owned context echoed back by the driver to the wrapper functions so the +// wrapper can recover which (hDriver, functionName) fired. Instances live for the +// life of the process in a tracing-layer registry (addresses must stay stable). +struct LoaderExtensionContext { + ze_driver_handle_t hDriver; + std::string functionName; +}; + typedef struct tracer_array_entry { zel_ze_all_callbacks_t corePrologues; zel_ze_all_callbacks_t coreEpilogues; zel_zer_all_callbacks_t runtimePrologues; zel_zer_all_callbacks_t runtimeEpilogues; void *pUserData; + // Per-tracer extension-function callbacks, copied by value into the active + // tracer array so the lock-free fan-out can walk them. + std::map extensionCallbacks; } tracer_array_entry_t; typedef struct tracerArray { @@ -60,6 +93,13 @@ struct APITracerImp : APITracer { ze_result_t resetAllCallbacks() override; ze_result_t enableTracer(ze_bool_t enable) override; + // Registers/clears one extension-function prologue or epilogue slot for + // (hDriver, functionName). Only valid while the tracer is disabled. + ze_result_t registerExtensionCallback(ze_driver_handle_t hDriver, + const char *functionName, + zel_tracer_reg_t callback_type, + zel_pfnDriverExtensionFunctionCb_t pCallback); + tracer_array_entry_t tracerFunctions; tracingState_t tracingState; @@ -268,4 +308,20 @@ APITracerWrapperImp(TFunction_pointer zeApiPtr, TParams paramsStruct, return ret; } +// Returns the stable, process-lifetime loader context for (hDriver, +// functionName), creating it on first use. The returned pointer is handed to the +// driver and echoed back to the wrapper functions below. +LoaderExtensionContext *getOrCreateLoaderExtensionContext(ze_driver_handle_t hDriver, + const char *functionName); + +// Loader-owned wrappers registered with the driver. The driver calls these from +// the body of the intercepted extension function; pLoaderContext is the +// LoaderExtensionContext* returned by getOrCreateLoaderExtensionContext. +void ZE_APICALL loaderExtensionPrologue(void *pParams, ze_result_t result, + void *pLoaderContext, + void **ppTracerInstanceUserData); +void ZE_APICALL loaderExtensionEpilogue(void *pParams, ze_result_t result, + void *pLoaderContext, + void **ppTracerInstanceUserData); + } // namespace tracing_layer diff --git a/source/layers/tracing/ze_tracing.cpp b/source/layers/tracing/ze_tracing.cpp index 1e8ec634..5a8ece38 100644 --- a/source/layers/tracing/ze_tracing.cpp +++ b/source/layers/tracing/ze_tracing.cpp @@ -6,6 +6,7 @@ */ #include "tracing.h" +#include "tracing_imp.h" #include "ze_tracing_layer.h" #include "layers/zel_tracing_api.h" #include "layers/zel_tracing_ddi.h" @@ -86,6 +87,61 @@ zelGetTracerApiProcAddrTable( return result; } +/////////////////////////////////////////////////////////////////////////////// +/// @brief Registers a prologue or epilogue callback on a tracer for a named +/// extension function of a specific driver. See loader/ze_loader.h. +ZE_DLLEXPORT ze_result_t ZE_APICALL +zelTracerDriverExtensionRegisterCallback( + zel_tracer_handle_t hTracer, + ze_driver_handle_t hDriver, + const char* functionName, + zel_tracer_reg_t callback_type, + zel_pfnDriverExtensionFunctionCb_t pCallback + ) +{ + if( nullptr == hTracer ) + return ZE_RESULT_ERROR_INVALID_NULL_HANDLE; + if( nullptr == hDriver ) + return ZE_RESULT_ERROR_INVALID_NULL_HANDLE; + if( nullptr == functionName ) + return ZE_RESULT_ERROR_INVALID_NULL_POINTER; + + // Record the app callback on the tracer (only valid while disabled). + ze_result_t result = tracing_layer::APITracer::fromHandle(hTracer) + ->registerExtensionCallback(hDriver, functionName, callback_type, pCallback); + if( result != ZE_RESULT_SUCCESS ) + return result; + + // Clearing a slot needs no driver interaction; the installed wrapper simply + // finds no registered tracer and no-ops. + if( nullptr == pCallback ) + return ZE_RESULT_SUCCESS; + + // Install (idempotently) the loader-owned wrapper on this driver so the + // driver invokes it from the body of the named extension function. Resolved + // by name via the downstream driver dispatch, exactly like the per-API + // GetExtensionFunctionAddress interceptor. + auto pfnGetExtensionFunctionAddress = + tracing_layer::context.zeDdiTable.Driver.pfnGetExtensionFunctionAddress; + if( nullptr == pfnGetExtensionFunctionAddress ) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + + void* pfnRaw = nullptr; + result = pfnGetExtensionFunctionAddress( + hDriver, "zelDriverSetLoaderCallbackForExtension", &pfnRaw ); + if( result != ZE_RESULT_SUCCESS ) + return result; + if( nullptr == pfnRaw ) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + + auto *ctx = tracing_layer::getOrCreateLoaderExtensionContext( hDriver, functionName ); + auto pfnSet = reinterpret_cast( pfnRaw ); + return pfnSet( hDriver, functionName, + &tracing_layer::loaderExtensionPrologue, + &tracing_layer::loaderExtensionEpilogue, + ctx ); +} + ZE_DLLEXPORT ze_result_t ZE_APICALL zelLoaderGetVersion(zel_component_version_t *version) { diff --git a/source/lib/ze_lib.cpp b/source/lib/ze_lib.cpp index e2aec2af..58f0a284 100644 --- a/source/lib/ze_lib.cpp +++ b/source/lib/ze_lib.cpp @@ -349,6 +349,20 @@ namespace ze_lib #endif isInitialized = true; } + + // Env-enabled tracing (ZE_ENABLE_TRACING_LAYER) never calls + // zelEnableTracingLayer, so propagate the extension-function tracing enable + // to each driver here. Drivers whose DDI tables were initialized via the + // proc-address-table setup above bypass init_driver (and its propagation), + // so this is the reliable point to open the driver-side tracing gate. +#ifndef L0_STATIC_LOADER_BUILD + if (ZE_RESULT_SUCCESS == result && loader::context && + loader::context->tracingLayerEnabled) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, true); + } + } +#endif #ifdef L0_STATIC_LOADER_BUILD std::call_once(ze_lib::context->initTeardownCallbacksOnce, [this]() { if (!delayContextDestruction) { @@ -722,64 +736,4 @@ zelDisableTracingLayer() return ZE_RESULT_SUCCESS; } -ze_result_t ZE_APICALL -zelDriverSetExtensionFunctionCallback( - ze_driver_handle_t hDriver, - const char* functionName, - void* pUserData, - zel_pfnDriverExtensionFunctionCb_t prologue, - zel_pfnDriverExtensionFunctionCb_t epilogue - ) -{ - if( nullptr == hDriver ) - return ZE_RESULT_ERROR_INVALID_NULL_HANDLE; - if( nullptr == functionName ) - return ZE_RESULT_ERROR_INVALID_NULL_POINTER; - if( ze_lib::destruction ) - return ZE_RESULT_ERROR_UNINITIALIZED; - - // Type of the driver-side registration entry point, discovered by name. - typedef ze_result_t (ZE_APICALL *zelDriverSetExtensionFunctionCallback_t)( - ze_driver_handle_t, const char*, void*, - zel_pfnDriverExtensionFunctionCb_t, zel_pfnDriverExtensionFunctionCb_t ); - - // Resolve the driver's registration entry via the standard extension-address - // lookup on this specific driver. The driver owns the registry and the - // invocation of the callbacks from inside the extension-function body. - auto pfnGetExtensionFunctionAddress = - ze_lib::context->zeDdiTable.load()->Driver.pfnGetExtensionFunctionAddress; - if( nullptr == pfnGetExtensionFunctionAddress ) { - if( !ze_lib::context->isInitialized ) - return ZE_RESULT_ERROR_UNINITIALIZED; - return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; - } - - void* pfnRaw = nullptr; - ze_result_t result = pfnGetExtensionFunctionAddress( - hDriver, "zelDriverSetExtensionFunctionCallback", &pfnRaw ); - if( result != ZE_RESULT_SUCCESS ) - return result; - if( nullptr == pfnRaw ) - return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; - - // Sync this driver's global extension-tracing gate to the current tracing - // state. Covers static ZE_ENABLE_TRACING_LAYER enablement (the driver becomes - // usable only by the time the app registers) and drivers registered on after - // a dynamic enable. Disabling is handled centrally by zelDisableTracingLayer - // (respecting sticky env), so only propagate the enable here. - bool tracingOn = ze_lib::context->tracingLayerEnableCounter.load() > 0; - if( !tracingOn && ze_lib::context->loaderContext ) - tracingOn = ze_lib::context->loaderContext->tracingLayerEnabled; - if( tracingOn ) { - void* pfnEnableRaw = nullptr; - if( ZE_RESULT_SUCCESS == pfnGetExtensionFunctionAddress( - hDriver, "zelDriverEnableTracing", &pfnEnableRaw ) && pfnEnableRaw ) { - reinterpret_cast(pfnEnableRaw)( hDriver, true ); - } - } - - auto pfnSet = reinterpret_cast( pfnRaw ); - return pfnSet( hDriver, functionName, pUserData, prologue, epilogue ); -} - } //extern "c" \ No newline at end of file diff --git a/source/lib/zel_tracing_libapi.cpp b/source/lib/zel_tracing_libapi.cpp index 8d3612d8..98d7218f 100644 --- a/source/lib/zel_tracing_libapi.cpp +++ b/source/lib/zel_tracing_libapi.cpp @@ -10,6 +10,7 @@ * Perhaps generate this from scripts in the future. */ #include "ze_lib.h" +#include "loader/ze_loader.h" extern "C" { @@ -175,4 +176,34 @@ zelTracerSetEnabled( return pfnSetEnabled( hTracer, enable ); } +/////////////////////////////////////////////////////////////////////////////// +/// @brief Registers a prologue/epilogue callback on a tracer for a named +/// extension function of a specific driver. See loader/ze_loader.h. +ze_result_t ZE_APICALL +zelTracerDriverExtensionRegisterCallback( + zel_tracer_handle_t hTracer, ///< [in] handle of the tracer + ze_driver_handle_t hDriver, ///< [in] handle of the driver instance + const char* functionName, ///< [in] extension function name to intercept + zel_tracer_reg_t callback_type, ///< [in] prologue or epilogue + zel_pfnDriverExtensionFunctionCb_t pCallback ///< [in][optional] handler (null clears slot) + ) +{ + if(ze_lib::destruction) + return ZE_RESULT_ERROR_UNINITIALIZED; + if(!ze_lib::context->tracing_lib) + return ZE_RESULT_ERROR_UNINITIALIZED; + + typedef ze_result_t (ZE_APICALL *ze_pfnRegisterExtCallback_t)( + zel_tracer_handle_t, ze_driver_handle_t, const char*, + zel_tracer_reg_t, zel_pfnDriverExtensionFunctionCb_t ); + + auto func = reinterpret_cast( + GET_FUNCTION_PTR(ze_lib::context->tracing_lib, + "zelTracerDriverExtensionRegisterCallback") ); + if(func) + return func( hTracer, hDriver, functionName, callback_type, pCallback ); + + return ZE_RESULT_ERROR_UNINITIALIZED; +} + } // extern "C" diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index 45394fd2..af924fff 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -265,7 +265,7 @@ set_property(TEST tests_tracing_layer_state_enabled_via_environment_and_dynamic add_test(NAME tests_tracing_layer_state_enabled_via_environment_disable_dynamic COMMAND tests --gtest_filter=*TracingLayerState.GivenTracingLayerEnabledViaEnvironmentAndDynamicallyWhenDisablingDynamicTracingThenStateRemainsTrue) set_property(TEST tests_tracing_layer_state_enabled_via_environment_disable_dynamic PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1") -# Extension-function callback (zelDriverSetExtensionFunctionCallback) tests +# Extension-function callback (zelTracerDriverExtensionRegisterCallback) tests # Dynamic control suite (toggles tracing at runtime). add_test(NAME tests_ext_fn_callback COMMAND tests --gtest_filter=*ExtFnCallback.*) set_property(TEST tests_ext_fn_callback PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1") diff --git a/test/loader_ext_fn_callback.cpp b/test/loader_ext_fn_callback.cpp index 001cdf37..d4026f0b 100644 --- a/test/loader_ext_fn_callback.cpp +++ b/test/loader_ext_fn_callback.cpp @@ -9,6 +9,7 @@ #include "gtest/gtest.h" #include "loader/ze_loader.h" +#include "layers/zel_tracing_api.h" #include "ze_api.h" #include @@ -21,7 +22,8 @@ typedef ze_result_t (ZE_APICALL *pfnSampleExtFunc_t)( constexpr uintptr_t kInstanceSentinel = 0xABCD1234u; -// State the prologue/epilogue callbacks record into, reached via pTracerUserData. +// State the prologue/epilogue callbacks record into, reached via pTracerUserData +// (the tracer's pUserData set at zelTracerCreate time). struct CallbackState { int prologCount = 0; int epilogCount = 0; @@ -71,14 +73,39 @@ pfnSampleExtFunc_t getSampleExtFunc(ze_driver_handle_t hDriver) { return reinterpret_cast(addr); } -void unregister(ze_driver_handle_t hDriver, const char* name) { - zelDriverSetExtensionFunctionCallback(hDriver, name, nullptr, nullptr, nullptr); +// Creates a disabled tracer whose pUserData is delivered to the callbacks. +zel_tracer_handle_t createTracer(void* pUserData) { + zel_tracer_desc_t desc = {}; + desc.stype = ZEL_STRUCTURE_TYPE_TRACER_DESC; + desc.pUserData = pUserData; + zel_tracer_handle_t hTracer = nullptr; + EXPECT_EQ(ZE_RESULT_SUCCESS, zelTracerCreate(&desc, &hTracer)); + EXPECT_NE(nullptr, hTracer); + return hTracer; +} + +// Registers both prologue and epilogue for a named extension function. +void registerCbs(zel_tracer_handle_t hTracer, ze_driver_handle_t hDriver, + const char* name) { + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, name, ZEL_REGISTER_PROLOGUE, prologueCb)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, name, ZEL_REGISTER_EPILOGUE, epilogueCb)); +} + +// Disables and destroys a tracer (destroy requires the disabled state). +void teardownTracer(zel_tracer_handle_t hTracer) { + EXPECT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, false)); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelTracerDestroy(hTracer)); } // --------------------------------------------------------------------------- // Dynamic control suite: tracing is toggled at runtime via // zelEnableTracingLayer/zelDisableTracingLayer (no ZE_ENABLE_TRACING_LAYER env). -// Each test balances enable/disable and unregisters so process state stays clean. +// Each test balances enable/disable and destroys its tracer so process state +// stays clean. // --------------------------------------------------------------------------- TEST(ExtFnCallback, PrologueAndEpilogueFireOnCall) { @@ -86,11 +113,12 @@ TEST(ExtFnCallback, PrologueAndEpilogueFireOnCall) { auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 21, &out)); @@ -104,18 +132,19 @@ TEST(ExtFnCallback, PrologueAndEpilogueFireOnCall) { EXPECT_EQ(ZE_RESULT_SUCCESS, state.epilogResult); EXPECT_EQ(kInstanceSentinel, state.instanceValueSeenInEpilog); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); } TEST(ExtFnCallback, RegisterBeforeFetchStillFires) { auto hDriver = getFirstDriver(); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); @@ -126,7 +155,7 @@ TEST(ExtFnCallback, RegisterBeforeFetchStillFires) { EXPECT_EQ(1, state.prologCount); EXPECT_EQ(1, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); } @@ -135,14 +164,19 @@ TEST(ExtFnCallback, UnregisterStopsCallbacks) { auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + // Clear both slots (null callback) while still disabled. EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, nullptr)); EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, nullptr, nullptr)); + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_EPILOGUE, nullptr)); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 7, &out)); @@ -151,17 +185,19 @@ TEST(ExtFnCallback, UnregisterStopsCallbacks) { EXPECT_EQ(0, state.prologCount); EXPECT_EQ(0, state.epilogCount); + teardownTracer(hTracer); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); } TEST(ExtFnCallback, UnknownFunctionNameNeverFires) { auto hDriver = getFirstDriver(); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback( - hDriver, "zeNeverImplementedExtFunc", &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeNeverImplementedExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); @@ -172,31 +208,55 @@ TEST(ExtFnCallback, UnknownFunctionNameNeverFires) { EXPECT_EQ(0, state.prologCount); EXPECT_EQ(0, state.epilogCount); - unregister(hDriver, "zeNeverImplementedExtFunc"); + teardownTracer(hTracer); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); } TEST(ExtFnCallback, NullArgumentsReturnErrors) { auto hDriver = getFirstDriver(); + CallbackState state; + auto hTracer = createTracer(&state); + + EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_HANDLE, + zelTracerDriverExtensionRegisterCallback( + nullptr, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, prologueCb)); EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_HANDLE, - zelDriverSetExtensionFunctionCallback(nullptr, "zeSampleExtFunc", - nullptr, prologueCb, epilogueCb)); + zelTracerDriverExtensionRegisterCallback( + hTracer, nullptr, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, prologueCb)); EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_POINTER, - zelDriverSetExtensionFunctionCallback(hDriver, nullptr, nullptr, - prologueCb, epilogueCb)); + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, nullptr, ZEL_REGISTER_PROLOGUE, prologueCb)); + + teardownTracer(hTracer); +} + +// Registration is only permitted while the tracer is disabled. +TEST(ExtFnCallback, RegisterWhileEnabledIsRejected) { + auto hDriver = getFirstDriver(); + + CallbackState state; + auto hTracer = createTracer(&state); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); + + EXPECT_EQ(ZE_RESULT_ERROR_INVALID_ARGUMENT, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, prologueCb)); + + teardownTracer(hTracer); } -// Two-level gate: registered but tracing layer NOT enabled -> must not fire. +// Two-level gate: registered + tracer enabled but tracing layer NOT enabled +// (driver gate closed) -> must not fire. TEST(ExtFnCallback, NotEnabledDoesNotFire) { auto hDriver = getFirstDriver(); auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 9, &out)); @@ -205,7 +265,7 @@ TEST(ExtFnCallback, NotEnabledDoesNotFire) { EXPECT_EQ(0, state.prologCount); // gate closed -> no callbacks EXPECT_EQ(0, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); } // Disabling the tracing layer stops callbacks even while still registered. @@ -214,11 +274,12 @@ TEST(ExtFnCallback, DisableStopsCallbacks) { auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 1, &out)); @@ -229,13 +290,48 @@ TEST(ExtFnCallback, DisableStopsCallbacks) { EXPECT_EQ(1, state.prologCount); // no additional fire after disable EXPECT_EQ(1, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); } -// Underflow guard (fix #1): a disable with no matching enable must be a safe -// no-op. If the unsigned counter had underflowed, the subsequent enable would -// not detect the 0->1 edge, the driver would never be enabled, and the callback -// would not fire - so a passing "fires" assertion proves no corruption occurred. +// Multiple tracers registered for the same function stack: all fire on one call. +// This is the capability the tracer-based design adds over the old per-driver +// last-writer-wins registry. +TEST(ExtFnCallback, MultipleTracersStack) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState s1; + CallbackState s2; + auto hTracer1 = createTracer(&s1); + auto hTracer2 = createTracer(&s2); + registerCbs(hTracer1, hDriver, "zeSampleExtFunc"); + registerCbs(hTracer2, hDriver, "zeSampleExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer1, true)); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer2, true)); + + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 11, &out)); + + EXPECT_EQ(22u, out); + EXPECT_EQ(1, s1.prologCount); + EXPECT_EQ(1, s1.epilogCount); + EXPECT_EQ(1, s2.prologCount); + EXPECT_EQ(1, s2.epilogCount); + EXPECT_EQ(kInstanceSentinel, s1.instanceValueSeenInEpilog); + EXPECT_EQ(kInstanceSentinel, s2.instanceValueSeenInEpilog); + + teardownTracer(hTracer1); + teardownTracer(hTracer2); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// Underflow guard: a disable with no matching enable must be a safe no-op. If +// the unsigned counter had underflowed, the subsequent enable would not detect +// the 0->1 edge, the driver gate would never open, and the callback would not +// fire - so a passing "fires" assertion proves no corruption occurred. TEST(ExtFnCallback, DisableWithoutEnableIsSafe) { auto hDriver = getFirstDriver(); auto fn = getSampleExtFunc(hDriver); @@ -246,26 +342,26 @@ TEST(ExtFnCallback, DisableWithoutEnableIsSafe) { EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); - ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 2, &out)); EXPECT_EQ(1, state.prologCount); // enable's 0->1 edge still worked EXPECT_EQ(1, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); } // --------------------------------------------------------------------------- // Environment suite: run with ZE_ENABLE_TRACING_LAYER=1 (separate ctest entry). -// Tracing is enabled statically at init; the app never calls -// zelEnableTracingLayer, and per documented behavior it stays enabled for the -// whole process. +// The driver gate is opened at init; the app never calls zelEnableTracingLayer, +// and per documented behavior it stays enabled for the whole process. // --------------------------------------------------------------------------- TEST(ExtFnCallbackEnviron, EnvKeepsTracingEnabled) { @@ -273,11 +369,11 @@ TEST(ExtFnCallbackEnviron, EnvKeepsTracingEnabled) { auto fn = getSampleExtFunc(hDriver); ASSERT_NE(nullptr, fn); - // No zelEnableTracingLayer call: the driver was enabled at init (Site A). + // No zelEnableTracingLayer call: the driver gate was opened at init. CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 4, &out)); @@ -286,7 +382,7 @@ TEST(ExtFnCallbackEnviron, EnvKeepsTracingEnabled) { EXPECT_EQ(1, state.prologCount); EXPECT_EQ(1, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); } // A spurious disable under static enablement must not turn tracing off (sticky @@ -299,9 +395,9 @@ TEST(ExtFnCallbackEnviron, EnvDisableIsNoOp) { EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); // counter 0 -> no-op CallbackState state; - EXPECT_EQ(ZE_RESULT_SUCCESS, - zelDriverSetExtensionFunctionCallback(hDriver, "zeSampleExtFunc", - &state, prologueCb, epilogueCb)); + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); uint32_t out = 0; EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 6, &out)); @@ -310,7 +406,7 @@ TEST(ExtFnCallbackEnviron, EnvDisableIsNoOp) { EXPECT_EQ(1, state.prologCount); // still fires: env-enabled tracing is sticky EXPECT_EQ(1, state.epilogCount); - unregister(hDriver, "zeSampleExtFunc"); + teardownTracer(hTracer); } } // namespace