diff --git a/include/loader/ze_loader.h b/include/loader/ze_loader.h index f71f5392..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 @@ -565,6 +566,106 @@ 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 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 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 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] 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] 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 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 +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 defined(__cplusplus) } // extern "C" #endif diff --git a/source/drivers/null/ze_null.cpp b/source/drivers/null/ze_null.cpp index 2418e033..1ba26151 100644 --- a/source/drivers/null/ze_null.cpp +++ b/source/drivers/null/ze_null.cpp @@ -47,6 +47,37 @@ 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, "zelDriverSetLoaderCallbackForExtension" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zelDriverSetLoaderCallbackForExtension ); + 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; + } + if( 0 == strcmp( name, "zelTestGetDriverTracingEnableCount" ) ) { + *ppFunctionAddress = reinterpret_cast( &driver::zelTestGetDriverTracingEnableCount ); + return ZE_RESULT_SUCCESS; + } + *ppFunctionAddress = nullptr; + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + }; + ////////////////////////////////////////////////////////////////////////// zeDdiTable.Device.pfnGet = []( ze_driver_handle_t, @@ -680,6 +711,104 @@ 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 the single loader wrapper registered for this function. + context_t::loader_extension_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: 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. + ze_sample_ext_func_params_t params = { &hDriver, &input, &pOutput }; + void* pInstanceData = nullptr; + ze_result_t result = ZE_RESULT_SUCCESS; + + 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.loaderEpilogue ) + cbs.loaderEpilogue( ¶ms, result, cbs.pLoaderContext, &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 ) + { + // Test hook: emulate a driver that advertises the "zelDriverEnableTracing" + // symbol but does not actually implement the capability. The loader's + // load-time probe invokes this with enable=false; returning UNSUPPORTED + // (without touching the gate) makes the loader treat this driver as + // unsupported and leave its gate permanently closed. + if( getenv_tobool( "ZEL_TEST_NULL_DRIVER_TRACING_UNSUPPORTED" ) ) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; + + if( enable ) + context.enableTracingTrueCount.fetch_add( 1 ); + context.extensionCallbacksEnabled.store( enable != 0 ); + return ZE_RESULT_SUCCESS; + } + + /////////////////////////////////////////////////////////////////////////// + /// @brief Test-only: report how many times the loader opened this driver's + /// extension-tracing gate (zelDriverEnableTracing with enable=true). + ze_result_t ZE_APICALL zelTestGetDriverTracingEnableCount( + ze_driver_handle_t /*hDriver*/, uint32_t* pCount ) + { + if( nullptr == pCount ) + return ZE_RESULT_ERROR_INVALID_NULL_POINTER; + *pCount = context.enableTracingTrueCount.load(); + return ZE_RESULT_SUCCESS; + } + + /////////////////////////////////////////////////////////////////////////// + /// @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 == loaderPrologue && nullptr == loaderEpilogue ) { + context.extensionCallbacks.erase( functionName ); + } else { + auto& entry = context.extensionCallbacks[ functionName ]; + entry.loaderPrologue = loaderPrologue; + entry.loaderEpilogue = loaderEpilogue; + entry.pLoaderContext = pLoaderContext; + } + 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); @@ -690,6 +819,20 @@ namespace driver context_t::~context_t() { + // Sever any link back into the loader/tracing layer before this driver + // goes away. The loader wrappers we were handed via + // zelDriverSetLoaderCallbackForExtension live in the tracing-layer .so, + // which may be unloaded around teardown. Close the gate and drop the + // stored wrapper pointers here, in our own destructor, so nothing can + // dereference them afterwards. This is done driver-side on purpose: the + // loader must not call into a driver during teardown (the driver's + // static state may already be gone), so each side cleans up its own. + { + std::lock_guard lock( extensionCallbackMutex ); + extensionCallbacksEnabled.store( false ); + extensionCallbacks.clear(); + } + for (auto handle : globalBaseNullHandle) { delete handle; diff --git a/source/drivers/null/ze_null.h b/source/drivers/null/ze_null.h index afff2802..c29da1bf 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,33 @@ namespace driver std::vector globalBaseNullHandle; bool ddiExtensionSupported = false; std::vector env_vars{}; + + // 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; + + // 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}; + + // Test observability: counts how many times the loader opened this + // driver's extension-tracing gate (zelDriverEnableTracing with + // enable=true). Exposed by name via "zelTestGetDriverTracingEnableCount" + // so tests can assert the loader skips the gate until a callback is + // registered (the lazy-gate optimization). + std::atomic enableTracingTrueCount{0}; + context_t(); ~context_t(); @@ -68,7 +100,44 @@ 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 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 + { + 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 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. + ze_result_t ZE_APICALL zelDriverEnableTracing( + ze_driver_handle_t hDriver, ze_bool_t enable ); + + // Test-only: returns the number of times zelDriverEnableTracing was called + // with enable=true (i.e. how many times the loader opened this driver's + // extension-tracing gate). Resolved by name "zelTestGetDriverTracingEnableCount". + ze_result_t ZE_APICALL zelTestGetDriverTracingEnableCount( + ze_driver_handle_t hDriver, uint32_t* pCount ); + extern context_t context; } // namespace driver diff --git a/source/layers/tracing/README.md b/source/layers/tracing/README.md index 3323bc81..3b5a5225 100644 --- a/source/layers/tracing/README.md +++ b/source/layers/tracing/README.md @@ -101,6 +101,56 @@ If the __callback_handler_function__ pointer is NULL, then no callback handler w These register callback functions can be called only when the __hTracer__ argument references a tracer that is in the disabled state. +### Registering callbacks for driver extension functions + +Extension functions retrieved by name via __zeDriverGetExtensionFunctionAddress__ return a raw driver function pointer that the application calls directly. These calls bypass the loader — and therefore the per-API tracing interceptors described above — so they cannot be traced with the core registration functions. To trace them, use: + +- __zelTracerDriverExtensionRegisterCallback(zel_tracer_handle_t hTracer, ze_driver_handle_t hDriver, const char\* functionName, zel_tracer_reg_t callback_type, zel_pfnDriverExtensionFunctionCb_t pCallback)__ + +This registers a prologue or epilogue handler on __hTracer__ for the extension function named __functionName__ on driver __hDriver__. It is declared in `include/loader/ze_loader.h`. + +Key points: +- Registration is keyed by the (__hDriver__, __functionName__) pair 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. +- `callback_type` is `ZEL_REGISTER_PROLOGUE` or `ZEL_REGISTER_EPILOGUE`; a null `pCallback` clears that slot. +- Like the core registration functions, this can be called only while __hTracer__ is in the disabled state. +- Multiple tracers may register the same function; their callbacks are stacked. +- It requires driver support (see **Driver Support** below) and returns `ZE_RESULT_ERROR_UNSUPPORTED_FEATURE` if the driver does not implement the required hooks. + +#### Callback signature + +Because an arbitrary extension function has no generated `..params_t` structure, the handler uses the generic signature `zel_pfnDriverExtensionFunctionCb_t` (in `include/loader/ze_loader.h`): + +``` +void (ZE_APICALL *zel_pfnDriverExtensionFunctionCb_t)( + void* pParams, // driver-defined parameter block (opaque; may be null) + ze_result_t result, // epilogue only: the function's return value + void* pTracerUserData, // per-tracer user data (from zelTracerCreate) + void** ppTracerInstanceUserData // per-call scratch for prologue->epilogue handoff +); +``` + +`pParams` points to a driver-defined layout for the named function; the driver documents its structure. The remaining parameters follow the same conventions as the core callback handlers described in **Callback Handlers**. + +#### When an extension callback fires + +A tracer's extension callback for a given (driver, function) fires on a call to that function only when **all** of the following hold: +1. The tracing layer is enabled (see **Enabling Tracing in the Loader**). +2. At least one tracer has registered a callback for that (driver, function), so the loader's wrapper is installed on the driver. +3. That specific tracer is enabled (via __zelTracerSetEnabled__) and registered the callback. + +These are the same layered semantics as core-API tracing: the tracing layer is the global switch, and each tracer must also be individually enabled. Registering a callback does not by itself cause it to fire; the tracer must be enabled. Disabling a tracer stops its extension callbacks from firing but leaves the registration in place. + +Conditions 1 and 2 together open the driver's extension-tracing gate, and they may be satisfied in either order: as an optimization the loader leaves the gate closed until the first extension callback is registered (so enabling the tracing layer with no extension callbacks costs nothing), then opens it on that first registration. Enabling the layer before or after registering a callback therefore produces the same result. + +#### Driver Support + +Extension-function tracing requires the driver to implement two hooks, discoverable by name through __zeDriverGetExtensionFunctionAddress__: + +- __zelDriverEnableTracing__ — a global gate the loader toggles when the tracing layer is enabled or disabled. While disabled, the driver must not invoke any registered wrapper. +- __zelDriverSetLoaderCallbackForExtension__ — installs (or clears) a single loader-owned prologue/epilogue wrapper for a named extension function. The driver invokes that wrapper around the body of the function. + +The loader probes these hooks at driver initialization. Drivers that do not implement them are treated as not supporting extension-function tracing, and __zelTracerDriverExtensionRegisterCallback__ returns `ZE_RESULT_ERROR_UNSUPPORTED_FEATURE`. The corresponding signatures (`zel_pfnDriverEnableTracing_t` and `zel_pfnDriverSetLoaderCallbackForExtension_t`) are defined in `include/loader/ze_loader.h`. + ## Reset All Callbacks __zelTracerResetAllCallbacks(zel_tracer_handle_t hTracer)__ can be used to set ALL prologue and epilogue callback handlers to NULL. @@ -117,7 +167,7 @@ Callback handlers are functions that are implemented by the application, and reg - __ppTracerInstanceUserData__ : a per-tracer, per-instance, per-thread storage location; typically used for passing data from the prologue to the epilogue. See example below. -## __ZeInit__ is traceable for all calls subsequent from the creation and enabling of the tracer itself. +## __zeInit__ is traceable for all calls subsequent from the creation and enabling of the tracer itself. ## Enabling, Disabling and Destruction The __tracer__ is created in a disabled state and must be explicitly enabled by calling __zelTracerSetEnabled__. The implementation guarantees that __prologue__ and __epilogue__ handlers for a given **L0 API** function will always be executed in pairs; i.e. @@ -262,4 +312,47 @@ void DynamicTracingExample( ... ) // Subsequent API calls will not be traced zeCommandListAppendLaunchKernel(hCommandList, hFunction, &launchArgs, nullptr, 0, nullptr); } + +// An example tracing a driver extension function obtained by name +void OnEnterMyExtFunc( + void* pParams, + ze_result_t result, + void* pTracerUserData, + void** ppTracerInstanceUserData ) +{ + // pParams points to the driver-defined parameter block for "zeMyExtFunc". + printf("entering zeMyExtFunc\n"); +} + +void ExtensionTracingExample( ze_driver_handle_t hDriver ) +{ + my_tracer_data_t tracer_data = {}; + zel_tracer_desc_t tracer_desc; + tracer_desc.stype = ZEL_STRUCTURE_TYPE_TRACER_DESC; + tracer_desc.pUserData = &tracer_data; + zel_tracer_handle_t hTracer; + zelTracerCreate(&tracer_desc, &hTracer); + + // Register a prologue for an extension function by name (tracer still disabled). + ze_result_t result = zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeMyExtFunc", ZEL_REGISTER_PROLOGUE, OnEnterMyExtFunc); + if (result == ZE_RESULT_ERROR_UNSUPPORTED_FEATURE) { + // The driver does not support extension-function tracing. + zelTracerDestroy(hTracer); + return; + } + + // The tracing layer must also be enabled for callbacks to fire. + zelEnableTracingLayer(); + zelTracerSetEnabled(hTracer, true); + + // Resolve and call the extension function directly; the prologue fires. + void* pfnRaw = nullptr; + zeDriverGetExtensionFunctionAddress(hDriver, "zeMyExtFunc", &pfnRaw); + // ... call the resolved function pointer as documented by the driver ... + + zelTracerSetEnabled(hTracer, false); + zelTracerDestroy(hTracer); + zelDisableTracingLayer(); +} ``` diff --git a/source/layers/tracing/tracing.h b/source/layers/tracing/tracing.h index d6c296c8..f0b3de7f 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,18 @@ 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; + // Records/clears one extension-function prologue or epilogue slot for + // (hDriver, functionName). pPrologueDelta / pEpilogueDelta (optional) each + // report whether this tracer's slot for that phase transitioned empty->active + // (+1), active->empty (-1), or was unchanged (0), so the caller can refcount + // the driver-side wrapper install/uninstall per phase. A single call changes + // at most one phase. + virtual ze_result_t registerExtensionCallback(ze_driver_handle_t hDriver, + const char *functionName, + zel_tracer_reg_t callback_type, + zel_pfnDriverExtensionFunctionCb_t pCallback, + int *pPrologueDelta = nullptr, + int *pEpilogueDelta = nullptr) = 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..b87f7695 100644 --- a/source/layers/tracing/tracing_imp.cpp +++ b/source/layers/tracing/tracing_imp.cpp @@ -7,6 +7,8 @@ #include "tracing_imp.h" +#include + namespace tracing_layer { thread_local ze_bool_t tracingInProgress = 0; @@ -118,6 +120,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 +129,52 @@ 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, int *pPrologueDelta, + int *pEpilogueDelta) { + + if (pPrologueDelta != nullptr) + *pPrologueDelta = 0; + if (pEpilogueDelta != nullptr) + *pEpilogueDelta = 0; + + // 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 &callbacks = this->tracerFunctions.extensionCallbacks; + + // Snapshot each phase's before/after presence so the caller can refcount the + // shared driver-side prologue and epilogue wrappers independently. + auto it = callbacks.find(key); + const bool hadPrologue = it != callbacks.end() && it->second.prologue != nullptr; + const bool hadEpilogue = it != callbacks.end() && it->second.epilogue != nullptr; + + auto &entry = callbacks[key]; + if (callback_type == ZEL_REGISTER_PROLOGUE) + entry.prologue = pCallback; + else + entry.epilogue = pCallback; + + const bool hasPrologue = entry.prologue != nullptr; + const bool hasEpilogue = entry.epilogue != nullptr; + + // Drop fully-cleared entries so the fan-out never carries dead keys. + if (!hasPrologue && !hasEpilogue) + callbacks.erase(key); + + if (pPrologueDelta != nullptr) + *pPrologueDelta = static_cast(hasPrologue) - static_cast(hadPrologue); + if (pEpilogueDelta != nullptr) + *pEpilogueDelta = static_cast(hasEpilogue) - static_cast(hadEpilogue); + + 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 +546,236 @@ 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. +// The per-phase refcounts track how many tracers currently hold a live prologue +// or epilogue for the key, gating that phase's driver-side wrapper install. +struct LoaderExtensionRegistryEntry { + LoaderExtensionContext ctx; + uint32_t prologueRefCount = 0; + uint32_t epilogueRefCount = 0; + LoaderExtensionRegistryEntry(ze_driver_handle_t hDriver, + const char *functionName) + : ctx(hDriver, functionName) {} +}; +std::mutex loaderExtensionContextMutex; +std::map loaderExtensionContexts; + +// Caller must hold loaderExtensionContextMutex. Returns the registry entry for +// the key, creating (and initializing its ctx) on first use. +LoaderExtensionRegistryEntry & +getOrCreateRegistryEntryLocked(const ExtensionFunctionKey &key, + ze_driver_handle_t hDriver, + const char *functionName) { + auto it = loaderExtensionContexts.find(key); + if (it == loaderExtensionContexts.end()) { + it = loaderExtensionContexts + .emplace(std::piecewise_construct, + std::forward_as_tuple(key), + std::forward_as_tuple(hDriver, functionName)) + .first; + } + return it->second; +} + +// Applies a signed delta to an unsigned refcount, clamping at zero. +uint32_t applyRefDelta(uint32_t count, int delta) { + if (delta > 0) + return count + static_cast(delta); + if (delta < 0) { + uint32_t dec = static_cast(-delta); + return dec > count ? 0u : count - dec; + } + return count; +} + +// 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); + return &getOrCreateRegistryEntryLocked(key, hDriver, functionName).ctx; +} + +LoaderExtensionInstallState +updateLoaderExtensionInstall(ze_driver_handle_t hDriver, + const char *functionName, int prologueDelta, + int epilogueDelta) { + ExtensionFunctionKey key{hDriver, functionName}; + std::lock_guard lock(loaderExtensionContextMutex); + auto &entry = getOrCreateRegistryEntryLocked(key, hDriver, functionName); + + const bool hadPrologue = entry.prologueRefCount > 0; + const bool hadEpilogue = entry.epilogueRefCount > 0; + entry.prologueRefCount = applyRefDelta(entry.prologueRefCount, prologueDelta); + entry.epilogueRefCount = applyRefDelta(entry.epilogueRefCount, epilogueDelta); + + LoaderExtensionInstallState state; + state.wantPrologue = entry.prologueRefCount > 0; + state.wantEpilogue = entry.epilogueRefCount > 0; + state.prologueInstallChanged = state.wantPrologue != hadPrologue; + state.epilogueInstallChanged = state.wantEpilogue != hadEpilogue; + state.ctx = &entry.ctx; + return state; +} + +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}; + + // If no epilogue trampoline is installed for this function, no epilogue will + // ever consume per-call instance data, so run the prologue callbacks inline + // with a throwaway instance slot and allocate nothing. This is the fast path + // when the app registered only prologues. + if (!ctx->epilogueInstalled.load(std::memory_order_relaxed)) { + 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() || + cbIt->second.prologue == nullptr) + continue; + void *instanceUserData = nullptr; + cbIt->second.prologue(pParams, result, tracerEntry.pUserData, + &instanceUserData); + } + } + pGlobalAPITracerContextImp->releaseActivetracersList(); + tracingInProgress = 0; + return; + } + + // An epilogue is installed: snapshot each participating tracer's epilogue + // callback and per-call instance slot into a frame the epilogue wrapper will + // consume, so prologue-set instance data reaches the matching epilogue. + 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; + + // A non-null instance handle is a frame the prologue wrapper built (both + // phases installed): run its snapshotted epilogue callbacks and free it. + auto *frame = + static_cast(*ppTracerInstanceUserData); + if (frame != nullptr) { + 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; + return; + } + + // No frame: the prologue trampoline is not installed (the app registered only + // epilogues) or no tracer participated. Self-gather the active tracers' epilogue + // callbacks and run them with a fresh (null) instance slot. + if (pLoaderContext == nullptr) + return; + if (tracingInProgress) + return; + tracingInProgress = 1; + + auto *ctx = static_cast(pLoaderContext); + ExtensionFunctionKey key{ctx->hDriver, ctx->functionName}; + + 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() || + cbIt->second.epilogue == nullptr) + continue; + void *instanceUserData = nullptr; + cbIt->second.epilogue(pParams, result, tracerEntry.pUserData, + &instanceUserData); + } + } + pGlobalAPITracerContextImp->releaseActivetracersList(); + tracingInProgress = 0; +} + } // namespace tracing_layer diff --git a/source/layers/tracing/tracing_imp.h b/source/layers/tracing/tracing_imp.h index c03be8be..3be4c5e5 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,50 @@ 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; + // True while the epilogue trampoline is installed on the driver for this key. + // Read (lock-free) by the prologue wrapper to decide whether to allocate a + // per-call frame that threads instance data to the epilogue; written under the + // registry mutex with careful ordering in ze_tracing.cpp so a concurrent call + // never builds a frame the driver will not hand back to an epilogue. + std::atomic epilogueInstalled{false}; + LoaderExtensionContext(ze_driver_handle_t driver, const char *name) + : hDriver(driver), functionName(name) {} +}; + 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 +101,15 @@ 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, + int *pPrologueDelta = nullptr, + int *pEpilogueDelta = nullptr) override; + tracer_array_entry_t tracerFunctions; tracingState_t tracingState; @@ -268,4 +318,42 @@ 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); + +// Result of applying per-phase reference deltas to a (hDriver, functionName) +// install entry: the desired install state for each phase, whether each phase +// just crossed its 0<->1 boundary (so the caller must re-issue the driver setter), +// and the stable loader context to hand the driver. +struct LoaderExtensionInstallState { + bool wantPrologue; + bool wantEpilogue; + bool prologueInstallChanged; + bool epilogueInstallChanged; + LoaderExtensionContext *ctx; +}; + +// Applies prologue/epilogue reference deltas to the install refcounts for +// (hDriver, functionName), counting how many tracers currently hold a live +// callback of each phase. The single shared driver-side wrapper installs a phase +// trampoline on that phase's 0->1 edge and removes it on 1->0. The registry entry +// is never erased, so any in-flight driver call keeps a valid ctx. +LoaderExtensionInstallState updateLoaderExtensionInstall(ze_driver_handle_t hDriver, + const char *functionName, + int prologueDelta, + int epilogueDelta); + +// 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..32436c59 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,90 @@ 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). The + // per-phase deltas report whether this tracer just began (+1) or stopped (-1) + // tracing the prologue and/or epilogue of this (hDriver, functionName), so the + // shared driver-side wrappers can be refcounted per phase across all tracers. + // A single call changes at most one phase. + int prologueDelta = 0; + int epilogueDelta = 0; + ze_result_t result = tracing_layer::APITracer::fromHandle(hTracer) + ->registerExtensionCallback(hDriver, functionName, callback_type, pCallback, + &prologueDelta, &epilogueDelta); + if( result != ZE_RESULT_SUCCESS ) + return result; + + // Apply the per-phase deltas to the shared driver-side install refcounts. + // Nothing to do on the driver unless a phase just crossed its 0<->1 boundary + // (e.g. a 2nd tracer registering the same phase is idempotent). + auto install = tracing_layer::updateLoaderExtensionInstall( + hDriver, functionName, prologueDelta, epilogueDelta ); + if( !install.prologueInstallChanged && !install.epilogueInstallChanged ) + return ZE_RESULT_SUCCESS; + + // Resolve the driver's by-name loader-callback setter (same mechanism as the + // per-API GetExtensionFunctionAddress interceptor). Needed for both install + // and uninstall. + 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 pfnSet = reinterpret_cast( pfnRaw ); + + // Install only the phase trampolines the app actually asked for: a null + // pointer tells the driver to skip that phase entirely (no call-through, no + // per-call frame allocation), which is the point of the split refcount. When + // both phases are gone the null/null pair clears the driver's registration. + zel_pfnDriverExtensionFunctionCb_t pfnPrologue = + install.wantPrologue ? &tracing_layer::loaderExtensionPrologue : nullptr; + zel_pfnDriverExtensionFunctionCb_t pfnEpilogue = + install.wantEpilogue ? &tracing_layer::loaderExtensionEpilogue : nullptr; + void* pLoaderContext = + ( install.wantPrologue || install.wantEpilogue ) ? install.ctx : nullptr; + + // Order the epilogue-installed hint (read lock-free by the prologue wrapper to + // decide whether to build an instance frame) against the driver update so a + // concurrent call never builds a frame the driver won't hand back to an + // epilogue: when removing the epilogue clear the hint first; when installing it + // set the hint only after the driver already has the epilogue trampoline. + if( install.epilogueInstallChanged && !install.wantEpilogue ) + install.ctx->epilogueInstalled.store( false, std::memory_order_relaxed ); + + result = pfnSet( hDriver, functionName, pfnPrologue, pfnEpilogue, pLoaderContext ); + + if( install.epilogueInstallChanged && install.wantEpilogue ) + install.ctx->epilogueInstalled.store( true, std::memory_order_relaxed ); + + return result; +} + 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 14e20511..41a4bdcc 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 @@ -344,6 +349,23 @@ 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. + // Skipped until the first extension callback is registered; a later + // registration opens the gate via zelLoaderTracingLayerRegisterExtensionCallback. +#ifndef L0_STATIC_LOADER_BUILD + if (ZE_RESULT_SUCCESS == result && loader::context && + loader::context->tracingLayerEnabled && + loader::context->anyExtensionCallbackRegistered.load()) { + 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) { @@ -635,6 +657,16 @@ 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, + // but only once an extension callback has been registered. The common + // case registers none, so this keeps the enable a plain DDI-table swap; a + // later registration opens the gate via + // zelLoaderTracingLayerRegisterExtensionCallback. + if (loader::context && loader::context->anyExtensionCallbackRegistered.load()) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, true); + } + } } #endif return ZE_RESULT_SUCCESS; @@ -684,11 +716,31 @@ 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. Nothing to close until + // an extension callback has been registered (the enable path skipped the + // gate while the latch was false). + if (loader::context && !loader::context->tracingLayerEnabled && + loader::context->anyExtensionCallbackRegistered.load()) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, false); + } + } } #endif return ZE_RESULT_SUCCESS; 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/lib/zel_tracing_libapi.cpp b/source/lib/zel_tracing_libapi.cpp index 8d3612d8..c367f493 100644 --- a/source/lib/zel_tracing_libapi.cpp +++ b/source/lib/zel_tracing_libapi.cpp @@ -10,6 +10,8 @@ * Perhaps generate this from scripts in the future. */ #include "ze_lib.h" +#include "loader/ze_loader.h" +#include "../loader/ze_loader_api.h" extern "C" { @@ -175,4 +177,54 @@ 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 ZE_RESULT_ERROR_UNINITIALIZED; + + ze_result_t result = func( hTracer, hDriver, functionName, callback_type, pCallback ); + if( result == ZE_RESULT_SUCCESS ) { + // Notify the loader that an extension callback now exists, so the + // tracing-layer enable/disable toggle paths resume propagating the + // per-driver gate (they skip it while no extension callback is + // registered). Handles install-after-enable ordering. The loader owns + // the driver-side gate, so this must reach loader code in both builds. + #ifdef L0_STATIC_LOADER_BUILD + if( ze_lib::context->loader ) { + typedef ze_result_t (ZE_APICALL *notify_t)(); + auto notify = reinterpret_cast( + GET_FUNCTION_PTR(ze_lib::context->loader, + "zelLoaderTracingLayerRegisterExtensionCallback") ); + if( notify ) + notify(); + } + #else + zelLoaderTracingLayerRegisterExtensionCallback(); + #endif + } + return result; +} + } // extern "C" diff --git a/source/loader/ze_loader.cpp b/source/loader/ze_loader.cpp index 14a0ef27..e0dfea90 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,68 @@ namespace loader return true; } + // Resolve, validate, and cache this driver's global extension-tracing gate + // hook exactly once. Normally this happens at driver-init time, but drivers + // that bypass init_driver (e.g. null-driver / proc-address-table setups) are + // resolved lazily on their first enable/disable toggle instead - hence the + // one-shot guard rather than a fixed call site. Resolving the pointer only + // proves the symbol exists - a driver may expose a stub that returns + // UNSUPPORTED - so we also probe it by invoking with enable=false. That probe + // is side-effect- + // free at first-init because "disabled" is the default state (the call sets + // the gate to the value it already has), yet its return value is the + // authoritative capability signal: a real implementation returns SUCCESS, a + // stub returns UNSUPPORTED (or another error). Only a SUCCESS probe caches the + // hook, so afterwards the cached result is authoritative and runtime toggles + // just null-check driver.pfnDriverEnableTracing - no by-name lookup, no + // repeated capability query. + // + // The one-shot guard is essential: resolution can be triggered more than once + // per driver (init_driver may run for both zeInitDrivers and zeInit, and the + // enable path may also call in), and because the probe writes the gate, + // re-running it after tracing was legitimately enabled would clobber the open + // gate back to disabled. Probing once - immediately followed by the paired + // enable-propagation at the call site - keeps the final state correct. + void resolveDriverExtensionTracingHook(driver_t &driver) { + if (driver.driverEnableTracingResolved) + return; + driver.driverEnableTracingResolved = true; + driver.pfnDriverEnableTracing = nullptr; + + auto pfnGetExtensionFunctionAddress = driver.dditable.ze.Driver.pfnGetExtensionFunctionAddress; + if (nullptr == pfnGetExtensionFunctionAddress) + return; + + void *pfnRaw = nullptr; + // Global driver-level hook; the handle is not needed to resolve it. + if (ZE_RESULT_SUCCESS != pfnGetExtensionFunctionAddress(nullptr, "zelDriverEnableTracing", &pfnRaw) || + nullptr == pfnRaw) + return; + + // Probe with the benign default (disabled). A driver that truly supports + // the gate returns SUCCESS; one that only advertises the symbol returns a + // non-SUCCESS result and stays unsupported (cached pointer left null). + auto pfnEnableTracing = reinterpret_cast(pfnRaw); + if (ZE_RESULT_SUCCESS == pfnEnableTracing(nullptr, false)) + driver.pfnDriverEnableTracing = pfnEnableTracing; + } + + ze_result_t enableDriverExtensionTracing(driver_t &driver, ze_bool_t enable) { + // Ensure the gate hook is resolved and capability-probed. This is a no-op + // after the first call (one-shot guard), but it is essential here because + // some drivers (e.g. null-driver / proc-address-table setups) initialize + // their DDI tables without ever going through init_driver, so this enable + // path is their only resolution point. After resolution a null pointer + // means the driver does not support extension-function tracing, keeping + // the enable/disable toggle a simple null check + call with no per-toggle + // lookup or query. + resolveDriverExtensionTracingHook(driver); + if (nullptr == driver.pfnDriverEnableTracing) + return ZE_RESULT_ERROR_UNSUPPORTED_FEATURE; // driver doesn't support it; skip + + return driver.pfnDriverEnableTracing(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 +659,26 @@ namespace loader return ZE_RESULT_ERROR_UNINITIALIZED; } + // Resolve this driver's extension-tracing gate hook once, now that its DDI + // table is available. Done here (rather than inside the DDI-init block + // above) so drivers that populate their dispatch tables during discovery + // and skip that block still get their hook resolved deterministically. + resolveDriverExtensionTracingHook(driver); + + // 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. + // Skipped entirely until the first extension callback is registered - the + // common case registers none, so a late driver need not toggle its gate. + if (anyExtensionCallbackRegistered.load() && + (tracingLayerEnabled || + (ze_lib::context && ze_lib::context->tracingLayerEnableCounter.load() > 0))) { + enableDriverExtensionTracing(driver, true); + } + return ZE_RESULT_SUCCESS; } diff --git a/source/loader/ze_loader_api.cpp b/source/loader/ze_loader_api.cpp index 514d6d31..4cfed837 100644 --- a/source/loader/ze_loader_api.cpp +++ b/source/loader/ze_loader_api.cpp @@ -9,6 +9,7 @@ */ #include "ze_loader_internal.h" +#include "../lib/ze_lib.h" #if defined(__cplusplus) extern "C" { @@ -59,6 +60,37 @@ zelLoaderGetContext() { return loader::context; } +/////////////////////////////////////////////////////////////////////////////// +/// @brief Notify the loader that an extension-function tracing callback was +/// registered (see ze_loader_api.h). +ZE_DLLEXPORT ze_result_t ZE_APICALL +zelLoaderTracingLayerRegisterExtensionCallback() { + if (nullptr == loader::context) + return ZE_RESULT_ERROR_UNINITIALIZED; + + // Monotonic latch: once any extension callback is registered, the toggle + // paths resume propagating the per-driver gate. Only the 0 -> 1 transition + // needs to open gates now; later registrations are cheap no-ops here. + bool wasRegistered = loader::context->anyExtensionCallbackRegistered.exchange(true); + if (wasRegistered) + return ZE_RESULT_SUCCESS; + + // First registration. If the tracing layer is already active (statically via + // ZE_ENABLE_TRACING_LAYER, or dynamically via zelEnableTracingLayer), the + // enable path already ran and skipped the per-driver loop because the latch + // was false, so open the driver-side gate now (install-after-enable + // ordering). Otherwise a subsequent enable opens it. This runs inside the + // loader binary, so ze_lib::context->tracingLayerEnableCounter reflects the + // authoritative dynamic-enable state for both the dynamic and static builds. + if (loader::context->tracingLayerEnabled || + (ze_lib::context && ze_lib::context->tracingLayerEnableCounter.load() > 0)) { + for (auto &drv : loader::context->zeDrivers) { + loader::enableDriverExtensionTracing(drv, true); + } + } + return ZE_RESULT_SUCCESS; +} + /////////////////////////////////////////////////////////////////////////////// /// @brief Internal function for Setting the ZE ddi table for the Tracing Layer. /// diff --git a/source/loader/ze_loader_api.h b/source/loader/ze_loader_api.h index bbf4c09f..01e2e7a2 100644 --- a/source/loader/ze_loader_api.h +++ b/source/loader/ze_loader_api.h @@ -88,6 +88,21 @@ zelLoaderTranslateHandleInternal( void **handleOut); //Output: Pointer to handleOut is set to driver handle if successful +/////////////////////////////////////////////////////////////////////////////// +/// @brief Notify the loader that an extension-function tracing callback was +/// registered. +/// +/// @details +/// - Sets the monotonic "any extension callback registered" latch so the +/// tracing-layer enable/disable toggle paths resume propagating the +/// per-driver extension-tracing gate (they skip it while the latch is +/// false). On the first registration, if the tracing layer is already +/// active, the driver-side gate is opened immediately (the enable path +/// skipped it while the latch was false - install-after-enable ordering). +ZE_DLLEXPORT ze_result_t ZE_APICALL +zelLoaderTracingLayerRegisterExtensionCallback(); + + #if defined(__cplusplus) } #endif \ No newline at end of file diff --git a/source/loader/ze_loader_internal.h b/source/loader/ze_loader_internal.h index 45dce29c..cc8f79a5 100644 --- a/source/loader/ze_loader_internal.h +++ b/source/loader/ze_loader_internal.h @@ -67,6 +67,16 @@ namespace loader ze_result_t zetddiInitResult = ZE_RESULT_ERROR_UNINITIALIZED; ze_result_t zesddiInitResult = ZE_RESULT_ERROR_UNINITIALIZED; ze_result_t zerddiInitResult = ZE_RESULT_ERROR_UNINITIALIZED; + // This driver's "zelDriverEnableTracing" gate hook, resolved once at + // driver-init time. A null pointer means the driver does not support + // extension-function tracing, so every runtime enable/disable toggle + // (which a perf tool like VTune calls frequently) is a simple null check + // plus a call - no by-name GetExtensionFunctionAddress lookup on the hot + // path. driverEnableTracingResolved guards the one-time capability probe: + // init_driver can run more than once per driver, and the probe call + // mutates the gate, so it must fire exactly once (before any real enable). + zel_pfnDriverEnableTracing_t pfnDriverEnableTracing = nullptr; + bool driverEnableTracingResolved = false; }; using driver_vector_t = std::vector< driver_t >; @@ -170,6 +180,14 @@ namespace loader bool debugTraceAdvanced = false; // true when ZE_ENABLE_LOADER_DEBUG_TRACE=2 or ZEL_ENABLE_LOADER_LOGGING=2 bool driverDDIPathDefault = false; bool tracingLayerEnabled = false; + // Monotonic latch set the first time any extension-function tracing + // callback is registered (via zelLoaderTracingLayerRegisterExtensionCallback). + // While it is false the tracing-layer enable/disable toggle paths skip the + // per-driver extension-tracing gate propagation entirely - the common case + // (e.g. VTune) registers no extension callbacks, so this keeps + // enable/disable a cheap DDI-table swap. Never reset (the optimization only + // needs to hold until the first registration). + std::atomic anyExtensionCallbackRegistered{false}; std::once_flag coreDriverSortOnce; std::once_flag sysmanDriverSortOnce; std::atomic sortingInProgress = {false}; @@ -184,4 +202,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..ed29f026 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,26 @@ 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 (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") + +# 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") + +# Unsupported-driver suite: driver advertises the gate hook but returns +# UNSUPPORTED when probed, so the loader must never open its gate. +add_test(NAME tests_ext_fn_callback_unsupported COMMAND tests --gtest_filter=*ExtFnCallbackUnsupported.*) +set_property(TEST tests_ext_fn_callback_unsupported PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=1;ZE_ENABLE_TRACING_LAYER=1;ZEL_TEST_NULL_DRIVER_TRACING_UNSUPPORTED=1") + +# Lazy-gate suite: the enable path must skip the per-driver gate until the first +# extension callback is registered, then open it lazily (install-after-enable). +# Own process so the monotonic "any callback registered" latch starts clear. +add_test(NAME tests_ext_fn_callback_lazy_gate COMMAND tests --gtest_filter=*ExtFnCallbackLazyGate.*) +set_property(TEST tests_ext_fn_callback_lazy_gate PROPERTY ENVIRONMENT "ZE_ENABLE_LOADER_DEBUG_TRACE=1;ZE_ENABLE_NULL_DRIVER=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 +515,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..28eca8e7 --- /dev/null +++ b/test/loader_ext_fn_callback.cpp @@ -0,0 +1,703 @@ +/* + * + * Copyright (C) 2026 Intel Corporation + * + * SPDX-License-Identifier: MIT + * + */ + +#include "gtest/gtest.h" + +#include "loader/ze_loader.h" +#include "layers/zel_tracing_api.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 +// (the tracer's pUserData set at zelTracerCreate time). +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); +} + +// 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 destroys its tracer so process state +// stays clean. +// --------------------------------------------------------------------------- + +TEST(ExtFnCallback, PrologueAndEpilogueFireOnCall) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + 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)); + + 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); + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, RegisterBeforeFetchStillFires) { + auto hDriver = getFirstDriver(); + + CallbackState state; + 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); + 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); + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, UnregisterStopsCallbacks) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + // Clear both slots (null callback) while still disabled. + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, nullptr)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + 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)); + + EXPECT_EQ(14u, out); + EXPECT_EQ(0, state.prologCount); + EXPECT_EQ(0, state.epilogCount); + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +TEST(ExtFnCallback, UnknownFunctionNameNeverFires) { + auto hDriver = getFirstDriver(); + + CallbackState state; + 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); + 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); + + 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, + zelTracerDriverExtensionRegisterCallback( + hTracer, nullptr, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, prologueCb)); + EXPECT_EQ(ZE_RESULT_ERROR_INVALID_NULL_POINTER, + 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 + 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; + 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)); + + EXPECT_EQ(18u, out); // body runs + EXPECT_EQ(0, state.prologCount); // gate closed -> no callbacks + EXPECT_EQ(0, state.epilogCount); + + teardownTracer(hTracer); +} + +// Disabling the tracing layer stops callbacks even while still registered. +TEST(ExtFnCallback, DisableStopsCallbacks) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + 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)); + 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); + + teardownTracer(hTracer); +} + +// 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()); +} + +// Exercises the driver-side install refcount that gates the single loader wrapper +// shared by every tracer on a (driver, function) pair. The refcount is a single +// count per (driver, function) - it counts distinct tracers holding at least one +// callback for it, not prologue/epilogue separately. Two tracers register the +// same function; when one unregisters (2->1) the wrapper must stay installed so +// the surviving tracer keeps firing, and only the last unregister (1->0) tears it +// down. The surviving tracer's callbacks are the observable proof that the shared +// wrapper was not removed early. +TEST(ExtFnCallback, MultipleTracersRefcount) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState s1; + CallbackState s2; + auto hTracer1 = createTracer(&s1); + auto hTracer2 = createTracer(&s2); + + // 0->1 then 1->2: both tracers reference the same function's single wrapper. + 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)); + + // Baseline: both tracers fire through the single shared wrapper. + 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); + + // 2->1: unregister tracer1 (registration requires the disabled state). The + // wrapper must remain installed because tracer2 still holds a reference. + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer1, false)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer1, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, nullptr)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer1, hDriver, "zeSampleExtFunc", ZEL_REGISTER_EPILOGUE, nullptr)); + + // The surviving tracer must still fire; tracer1 must not. This is the key + // regression check: a broken refcount that uninstalled at 2->1 would silently + // stop tracer2 from firing (the driver would no longer see the wrapper). + out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 12, &out)); + EXPECT_EQ(24u, out); + EXPECT_EQ(1, s1.prologCount); // unchanged - tracer1 unregistered + EXPECT_EQ(1, s1.epilogCount); + EXPECT_EQ(2, s2.prologCount); // fired again through the still-installed wrapper + EXPECT_EQ(2, s2.epilogCount); + + // 1->0: unregister the last tracer - the wrapper is now uninstalled and no + // callback fires for either tracer. + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer2, false)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer2, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, nullptr)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer2, hDriver, "zeSampleExtFunc", ZEL_REGISTER_EPILOGUE, nullptr)); + + out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 13, &out)); + EXPECT_EQ(26u, out); + EXPECT_EQ(1, s1.prologCount); + EXPECT_EQ(2, s2.prologCount); + + teardownTracer(hTracer1); + teardownTracer(hTracer2); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// Prologue-only install: the app registers just a prologue, so only the prologue +// trampoline is installed on the driver (the epilogue slot stays null). The +// prologue must fire and the epilogue must never fire. This exercises the +// split-refcount fast path where no per-call instance frame is built because no +// epilogue is installed to consume it. +TEST(ExtFnCallback, PrologueOnlyFires) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + auto hTracer = createTracer(&state); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, + prologueCb)); + + 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, 8, &out)); + + EXPECT_EQ(16u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(0, state.epilogCount); // no epilogue installed + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// Epilogue-only install: the app registers just an epilogue, so only the epilogue +// trampoline is installed on the driver (the prologue slot stays null). The +// epilogue must fire. Because no prologue ran, there is no instance frame, so the +// epilogue self-gathers from the active tracers and sees a null instance value. +TEST(ExtFnCallback, EpilogueOnlyFires) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + auto hTracer = createTracer(&state); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracer, hDriver, "zeSampleExtFunc", ZEL_REGISTER_EPILOGUE, + epilogueCb)); + + 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, 9, &out)); + + EXPECT_EQ(18u, out); + EXPECT_EQ(0, state.prologCount); // no prologue installed + EXPECT_EQ(1, state.epilogCount); + EXPECT_EQ(ZE_RESULT_SUCCESS, state.epilogResult); + EXPECT_EQ(0u, state.instanceValueSeenInEpilog); // no frame -> null instance + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// Per-phase refcount independence: one tracer supplies only the prologue and a +// different tracer supplies only the epilogue. Both phase trampolines end up +// installed on the driver from different sources, and both callbacks fire on a +// single call. Removing the prologue tracer's callback (prologue 1->0) must not +// disturb the epilogue tracer, whose epilogue keeps firing via the self-gather +// path - proving the two install refcounts are tracked independently. +TEST(ExtFnCallback, PerPhaseRefcountAcrossTracers) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState sProlog; + CallbackState sEpilog; + auto hTracerProlog = createTracer(&sProlog); + auto hTracerEpilog = createTracer(&sEpilog); + + // Prologue phase 0->1 from tracerProlog; epilogue phase 0->1 from tracerEpilog. + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracerProlog, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, + prologueCb)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracerEpilog, hDriver, "zeSampleExtFunc", ZEL_REGISTER_EPILOGUE, + epilogueCb)); + + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracerProlog, true)); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracerEpilog, true)); + + // Both phases installed (from different tracers): each fires once. + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 10, &out)); + EXPECT_EQ(20u, out); + EXPECT_EQ(1, sProlog.prologCount); + EXPECT_EQ(0, sProlog.epilogCount); + EXPECT_EQ(0, sEpilog.prologCount); + EXPECT_EQ(1, sEpilog.epilogCount); + + // Prologue 1->0: remove the prologue tracer's callback. The epilogue refcount + // is untouched, so the epilogue trampoline must remain installed. + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracerProlog, false)); + EXPECT_EQ(ZE_RESULT_SUCCESS, + zelTracerDriverExtensionRegisterCallback( + hTracerProlog, hDriver, "zeSampleExtFunc", ZEL_REGISTER_PROLOGUE, + nullptr)); + + out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 15, &out)); + EXPECT_EQ(30u, out); + EXPECT_EQ(1, sProlog.prologCount); // unchanged - prologue uninstalled + EXPECT_EQ(2, sEpilog.epilogCount); // epilogue still installed and firing + + teardownTracer(hTracerProlog); + teardownTracer(hTracerEpilog); + 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); + 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()); + + CallbackState state; + 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); + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +// --------------------------------------------------------------------------- +// Environment suite: run with ZE_ENABLE_TRACING_LAYER=1 (separate ctest entry). +// 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) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + // No zelEnableTracingLayer call: the driver gate was opened at init. + CallbackState state; + 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)); + + EXPECT_EQ(8u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(1, state.epilogCount); + + teardownTracer(hTracer); +} + +// 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; + 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)); + + EXPECT_EQ(12u, out); + EXPECT_EQ(1, state.prologCount); // still fires: env-enabled tracing is sticky + EXPECT_EQ(1, state.epilogCount); + + teardownTracer(hTracer); +} + +// --------------------------------------------------------------------------- +// Unsupported-driver suite: run with ZE_ENABLE_TRACING_LAYER=1 AND +// ZEL_TEST_NULL_DRIVER_TRACING_UNSUPPORTED=1 (separate ctest entry). The driver +// advertises "zelDriverEnableTracing" (the symbol resolves), but its +// implementation returns ZE_RESULT_ERROR_UNSUPPORTED_FEATURE - i.e. presence of +// the hook does not imply support. The loader's load-time capability probe must +// therefore treat the driver as unsupported (leaving its cached gate hook null) +// and never open the gate, so extension callbacks never fire even though the +// tracing layer is enabled and callbacks are registered. The extension function +// itself still works normally. +// --------------------------------------------------------------------------- + +TEST(ExtFnCallbackUnsupported, StubDriverNeverFiresCallbacks) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + + CallbackState state; + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); + + uint32_t out = 0; + // The extension function still succeeds and performs its work... + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 5, &out)); + EXPECT_EQ(10u, out); + + // ...but because the driver reported the gate hook as unsupported, the probe + // left the cached hook null and the gate was never opened, so no prologue or + // epilogue ever fired. + EXPECT_EQ(0, state.prologCount); + EXPECT_EQ(0, state.epilogCount); + + teardownTracer(hTracer); +} + +// --------------------------------------------------------------------------- +// Lazy-gate suite: proves the tracing-layer enable path skips the per-driver +// extension-tracing gate while no extension callback is registered, and opens +// it lazily on the first registration (install-after-enable ordering). Runs as +// its own ctest process so the process-wide monotonic "any callback registered" +// latch starts clear - no other ExtFnCallback test may share this process. +// --------------------------------------------------------------------------- + +typedef ze_result_t (ZE_APICALL *pfnGetEnableCount_t)(ze_driver_handle_t, uint32_t*); + +pfnGetEnableCount_t getEnableCountFunc(ze_driver_handle_t hDriver) { + void* addr = nullptr; + EXPECT_EQ(ZE_RESULT_SUCCESS, + zeDriverGetExtensionFunctionAddress( + hDriver, "zelTestGetDriverTracingEnableCount", &addr)); + EXPECT_NE(nullptr, addr); + return reinterpret_cast(addr); +} + +TEST(ExtFnCallbackLazyGate, LayerEnableSkipsGateUntilRegister) { + auto hDriver = getFirstDriver(); + auto fn = getSampleExtFunc(hDriver); + ASSERT_NE(nullptr, fn); + auto getCount = getEnableCountFunc(hDriver); + + // Fresh process: the loader has never opened this driver's gate. + uint32_t count = 999; + ASSERT_EQ(ZE_RESULT_SUCCESS, getCount(hDriver, &count)); + ASSERT_EQ(0u, count); + + // Enable the tracing layer with NO extension callback registered. The + // optimization must skip the per-driver enable loop, so the driver's gate is + // never toggled (without it, the loop would open the gate here -> count 1). + ASSERT_EQ(ZE_RESULT_SUCCESS, zelEnableTracingLayer()); + ASSERT_EQ(ZE_RESULT_SUCCESS, getCount(hDriver, &count)); + EXPECT_EQ(0u, count) << "enable loop ran despite no registered ext callback"; + + // Register a callback AFTER the layer was enabled. The install-after-enable + // path must now open the driver gate (0->1 latch transition) exactly once. + CallbackState state; + auto hTracer = createTracer(&state); + registerCbs(hTracer, hDriver, "zeSampleExtFunc"); + ASSERT_EQ(ZE_RESULT_SUCCESS, getCount(hDriver, &count)); + EXPECT_EQ(1u, count) << "registration did not open the gate after enable"; + + // End-to-end: the callback fires even though it was registered after enable. + ASSERT_EQ(ZE_RESULT_SUCCESS, zelTracerSetEnabled(hTracer, true)); + uint32_t out = 0; + EXPECT_EQ(ZE_RESULT_SUCCESS, fn(hDriver, 6, &out)); + EXPECT_EQ(12u, out); + EXPECT_EQ(1, state.prologCount); + EXPECT_EQ(1, state.epilogCount); + + teardownTracer(hTracer); + EXPECT_EQ(ZE_RESULT_SUCCESS, zelDisableTracingLayer()); +} + +} // namespace