From d1de8a8fcb93df4c092e98086808d9364b3e9d73 Mon Sep 17 00:00:00 2001 From: Manasjyoti Sharma Date: Fri, 3 Apr 2026 14:45:04 +0530 Subject: [PATCH] Fix stale deferred findings leak in LangChain cache wrappers and depth counter leak in LlamaIndex streaming wrappers LangChain: base_chat_model_generate_with_cache_wrapper and base_llm_generate_helper_wrapper were missing discard_deferred_findings(), allowing stale safety findings from a prior request on the same thread to leak into cached/helper response paths. Their non-cached counterparts already had this call. LlamaIndex: All 4 streaming wrappers (llm_stream_chat_wrapper, llm_astream_chat_wrapper, llm_stream_complete_wrapper, llm_astream_complete_wrapper) had _enter_safety() calls without guaranteed _exit_safety() on all code paths. If wrapped(), await, or LlamaIndexStreamingSafety() threw, the depth counter stayed elevated permanently, breaking safety for all subsequent requests on the thread. Restructured all 4 to use try/finally. --- .../instrumentation/langchain/safety.py | 4 ++ .../instrumentation/llamaindex/safety.py | 51 ++++++++----------- 2 files changed, 26 insertions(+), 29 deletions(-) diff --git a/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py index 3e1def5026..6d48c755bf 100644 --- a/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py +++ b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py @@ -70,6 +70,8 @@ def _safety_on_worker(): def base_chat_model_generate_with_cache_wrapper(wrapped, instance, args, kwargs): + from opentelemetry.instrumentation.fortifyroot import discard_deferred_findings + discard_deferred_findings() # FR: prevent stale findings from prior request on same thread response = wrapped(*args, **kwargs) _apply_chat_result_completion_safety(instance, response) return response @@ -91,6 +93,8 @@ def _completion_safety_on_worker(): def base_llm_generate_helper_wrapper(wrapped, instance, args, kwargs): + from opentelemetry.instrumentation.fortifyroot import discard_deferred_findings + discard_deferred_findings() # FR: prevent stale findings from prior request on same thread response = wrapped(*args, **kwargs) _apply_llm_result_completion_safety(instance, response) return response diff --git a/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py index afdadbafbe..00d506af91 100644 --- a/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py +++ b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py @@ -153,16 +153,12 @@ def llm_stream_chat_wrapper(wrapped, instance, args, kwargs): # FR: streaming s _enter_safety() try: updated_args, updated_kwargs = _apply_chat_prompt_safety(instance, args, kwargs) - except Exception: + span = trace.get_current_span() + span_name = f"{instance.__class__.__name__}.stream_chat" + safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.CHAT.value) + return wrap_stream(wrapped(*updated_args, **updated_kwargs), safety) + finally: _exit_safety() - raise - # NOTE: _exit_safety NOT called here — the generator from wrap_stream may - # outlive this call frame. The depth counter stays +1 during iteration, - # which is harmless (prevents spurious clears during nested calls). - span = trace.get_current_span() - span_name = f"{instance.__class__.__name__}.stream_chat" - safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.CHAT.value) - return wrap_stream(wrapped(*updated_args, **updated_kwargs), safety) async def llm_astream_chat_wrapper(wrapped, instance, args, kwargs): # FR: streaming safety @@ -183,14 +179,13 @@ def _safety_on_worker(): (updated_args, updated_kwargs), worker_findings = await asyncio.to_thread(_safety_on_worker) inject_deferred_findings(worker_findings) - except Exception: + result = wrapped(*updated_args, **updated_kwargs) + if _inspect.iscoroutine(result): # coroutine returning async gen (standard LI pattern) + result = await result + safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.CHAT.value) + return make_async_stream(result, safety) + finally: _exit_safety() - raise - result = wrapped(*updated_args, **updated_kwargs) - if _inspect.iscoroutine(result): # coroutine returning async gen (standard LI pattern) - result = await result - safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.CHAT.value) - return make_async_stream(result, safety) def llm_stream_complete_wrapper(wrapped, instance, args, kwargs): # FR: streaming safety @@ -201,13 +196,12 @@ def llm_stream_complete_wrapper(wrapped, instance, args, kwargs): # FR: streami _enter_safety() try: updated_args, updated_kwargs = _apply_completion_prompt_safety(instance, args, kwargs) - except Exception: + span = trace.get_current_span() + span_name = f"{instance.__class__.__name__}.stream_complete" + safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.COMPLETION.value) + return wrap_stream(wrapped(*updated_args, **updated_kwargs), safety) + finally: _exit_safety() - raise - span = trace.get_current_span() - span_name = f"{instance.__class__.__name__}.stream_complete" - safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.COMPLETION.value) - return wrap_stream(wrapped(*updated_args, **updated_kwargs), safety) async def llm_astream_complete_wrapper(wrapped, instance, args, kwargs): # FR: streaming safety @@ -228,14 +222,13 @@ def _safety_on_worker(): (updated_args, updated_kwargs), worker_findings = await asyncio.to_thread(_safety_on_worker) inject_deferred_findings(worker_findings) - except Exception: + result = wrapped(*updated_args, **updated_kwargs) + if _inspect.iscoroutine(result): + result = await result + safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.COMPLETION.value) + return make_async_stream(result, safety) + finally: _exit_safety() - raise - result = wrapped(*updated_args, **updated_kwargs) - if _inspect.iscoroutine(result): - result = await result - safety = LlamaIndexStreamingSafety(span, span_name, LLMRequestTypeValues.COMPLETION.value) - return make_async_stream(result, safety) _METHOD_WRAPPERS = {