From c964f1a49eab4e2d9d8e7fe7e0909d4cab53bba4 Mon Sep 17 00:00:00 2001 From: Manasjyoti Sharma Date: Tue, 17 Mar 2026 20:12:42 +0530 Subject: [PATCH] Implement phase 1H non-stream safety parity --- .../instrumentation/alephalpha/__init__.py | 6 + .../instrumentation/alephalpha/safety.py | 178 +++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 143 +++++ .../instrumentation/bedrock/__init__.py | 11 +- .../instrumentation/bedrock/safety.py | 434 ++++++++++++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 224 ++++++++ .../instrumentation/groq/__init__.py | 8 + .../instrumentation/groq/safety.py | 205 ++++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 189 +++++++ .../instrumentation/langchain/__init__.py | 6 + .../instrumentation/langchain/safety.py | 488 ++++++++++++++++++ .../pyproject.toml | 7 + .../tests/test_safety_hooks.py | 342 ++++++++++++ .../README.md | 3 + .../instrumentation/litellm/__init__.py | 368 +++++++++++++ .../instrumentation/litellm/safety.py | 445 ++++++++++++++++ .../instrumentation/litellm/version.py | 1 + .../project.json | 77 +++ .../pyproject.toml | 81 +++ .../tests/__init__.py | 1 + .../tests/test_safety_hooks.py | 481 +++++++++++++++++ .../llamaindex/dispatcher_wrapper.py | 19 + .../instrumentation/llamaindex/safety.py | 449 ++++++++++++++++ .../pyproject.toml | 7 + .../tests/test_safety_hooks.py | 388 ++++++++++++++ .../instrumentation/mistralai/__init__.py | 8 + .../instrumentation/mistralai/safety.py | 182 +++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 146 ++++++ .../instrumentation/ollama/__init__.py | 8 + .../instrumentation/ollama/safety.py | 229 ++++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 137 +++++ .../pyproject.toml | 5 + .../instrumentation/replicate/__init__.py | 6 + .../instrumentation/replicate/safety.py | 195 +++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 192 +++++++ .../instrumentation/sagemaker/__init__.py | 14 +- .../instrumentation/sagemaker/safety.py | 193 +++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 130 +++++ .../instrumentation/together/__init__.py | 6 + .../instrumentation/together/safety.py | 249 +++++++++ .../pyproject.toml | 9 + .../pytest.ini | 2 + .../tests/test_safety_hooks.py | 192 +++++++ .../instrumentation/transformers/safety.py | 198 +++++++ .../text_generation_pipeline_wrapper.py | 6 + .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 131 +++++ .../instrumentation/vertexai/__init__.py | 8 + .../instrumentation/vertexai/safety.py | 200 +++++++ .../pyproject.toml | 7 + .../tests/test_safety_hooks.py | 125 +++++ .../instrumentation/watsonx/__init__.py | 6 + .../instrumentation/watsonx/safety.py | 145 ++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 116 +++++ .../instrumentation/writer/__init__.py | 10 + .../instrumentation/writer/safety.py | 247 +++++++++ .../pyproject.toml | 9 + .../tests/test_safety_hooks.py | 192 +++++++ .../opentelemetry/semconv_ai/__init__.py | 3 + packages/traceloop-sdk/pyproject.toml | 7 + .../tests/test_litellm_instrumentation.py | 42 ++ .../tests/test_sdk_initialization.py | 1 + .../traceloop/sdk/instruments.py | 1 + .../traceloop/sdk/tracing/tracing.py | 17 + 72 files changed, 8014 insertions(+), 2 deletions(-) create mode 100644 packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/safety.py create mode 100644 packages/opentelemetry-instrumentation-alephalpha/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/safety.py create mode 100644 packages/opentelemetry-instrumentation-bedrock/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/safety.py create mode 100644 packages/opentelemetry-instrumentation-groq/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py create mode 100644 packages/opentelemetry-instrumentation-langchain/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-litellm/README.md create mode 100644 packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/__init__.py create mode 100644 packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/safety.py create mode 100644 packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/version.py create mode 100644 packages/opentelemetry-instrumentation-litellm/project.json create mode 100644 packages/opentelemetry-instrumentation-litellm/pyproject.toml create mode 100644 packages/opentelemetry-instrumentation-litellm/tests/__init__.py create mode 100644 packages/opentelemetry-instrumentation-litellm/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py create mode 100644 packages/opentelemetry-instrumentation-llamaindex/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/safety.py create mode 100644 packages/opentelemetry-instrumentation-mistralai/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/safety.py create mode 100644 packages/opentelemetry-instrumentation-ollama/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/safety.py create mode 100644 packages/opentelemetry-instrumentation-replicate/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/safety.py create mode 100644 packages/opentelemetry-instrumentation-sagemaker/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/safety.py create mode 100644 packages/opentelemetry-instrumentation-together/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/safety.py create mode 100644 packages/opentelemetry-instrumentation-transformers/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/safety.py create mode 100644 packages/opentelemetry-instrumentation-vertexai/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/safety.py create mode 100644 packages/opentelemetry-instrumentation-watsonx/tests/test_safety_hooks.py create mode 100644 packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/safety.py create mode 100644 packages/opentelemetry-instrumentation-writer/tests/test_safety_hooks.py create mode 100644 packages/traceloop-sdk/tests/test_litellm_instrumentation.py diff --git a/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/__init__.py b/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/__init__.py index c05d50f786..32f560e03f 100644 --- a/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/__init__.py +++ b/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/__init__.py @@ -16,6 +16,10 @@ set_completion_attributes, set_prompt_attributes, ) +from opentelemetry.instrumentation.alephalpha.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.alephalpha.utils import dont_throw from opentelemetry.instrumentation.alephalpha.version import __version__ from opentelemetry.instrumentation.instrumentor import BaseInstrumentor @@ -164,12 +168,14 @@ def _wrap( SpanAttributes.LLM_REQUEST_TYPE: llm_request_type.value, }, ) + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) input_event = _parse_prompt_event(args, kwargs) _handle_message_event(input_event, span, event_logger, kwargs) response = wrapped(*args, **kwargs) if response: + _apply_completion_safety(span, response, name) response_event = _parse_completion_event(response) _handle_completion_event(response_event, span, event_logger, response) if span.is_recording(): diff --git a/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/safety.py b/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/safety.py new file mode 100644 index 0000000000..4a08d753c4 --- /dev/null +++ b/packages/opentelemetry-instrumentation-alephalpha/opentelemetry/instrumentation/alephalpha/safety.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +from aleph_alpha_client import Prompt as AlephAlphaPrompt +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "AlephAlpha" + + +def _apply_prompt_safety(span, args, kwargs, span_name): + try: + request = kwargs.get("request") or (args[0] if args else None) + prompt = get_object_value(request, "prompt") + updated_prompt, changed = _mask_prompt_value( + span, + prompt, + span_name=span_name, + ) + if not changed: + return args, kwargs + + set_object_value(request, "prompt", updated_prompt) + return args, kwargs + except Exception: + return args, kwargs + + +def _apply_completion_safety(span, response, span_name): + try: + completions = get_object_value(response, "completions") + if not completions: + return + for index, completion_item in enumerate(completions): + completion = get_object_value(completion_item, "completion") + if not isinstance(completion, str): + continue + updated_completion, changed = _mask_completion_text( + span, + completion, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(completion_item, "completion", updated_completion) + except Exception: + return + + +def _mask_prompt_value(span, prompt, *, span_name): + if isinstance(prompt, str): + updated_text, changed = _mask_prompt_text( + span, + prompt, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if not changed: + return prompt, False + return AlephAlphaPrompt.from_text(updated_text), True + + if prompt is None: + return prompt, False + + prompt_json = prompt.to_json() if hasattr(prompt, "to_json") else prompt + updated_json, changed = _mask_prompt_json(span, prompt_json, span_name=span_name) + if not changed: + return prompt, False + + if isinstance(updated_json, list): + return AlephAlphaPrompt.from_json(updated_json), True + return updated_json, True + + +def _mask_prompt_json(span, prompt_json, *, span_name): + if isinstance(prompt_json, list): + updated_prompt = prompt_json + for index, item in enumerate(prompt_json): + updated_item, changed = _mask_prompt_item( + span, + item, + span_name=span_name, + segment_index=index, + ) + if not changed: + continue + if updated_prompt is prompt_json: + updated_prompt = list(prompt_json) + updated_prompt[index] = updated_item + return updated_prompt, updated_prompt is not prompt_json + + if isinstance(prompt_json, dict): + return _mask_prompt_item( + span, + prompt_json, + span_name=span_name, + segment_index=0, + ) + + return prompt_json, False + + +def _mask_prompt_item(span, item, *, span_name, segment_index): + if not isinstance(item, dict): + return item, False + + text_key = None + if item.get("type") in (None, "text") and isinstance(item.get("data"), str): + text_key = "data" + elif item.get("type") in (None, "text") and isinstance(item.get("text"), str): + text_key = "text" + + if text_key is None: + return item, False + + updated_text, changed = _mask_prompt_text( + span, + item[text_key], + span_name=span_name, + segment_index=segment_index, + segment_role="user", + ) + if not changed: + return item, False + + updated_item = dict(item) + updated_item[text_key] = updated_text + return updated_item, True + + +def _mask_prompt_text( + span, + text, + *, + span_name, + segment_index, + segment_role, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-alephalpha/pyproject.toml b/packages/opentelemetry-instrumentation-alephalpha/pyproject.toml index 69a1ccc9ff..c209109a46 100644 --- a/packages/opentelemetry-instrumentation-alephalpha/pyproject.toml +++ b/packages/opentelemetry-instrumentation-alephalpha/pyproject.toml @@ -12,6 +12,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -72,5 +73,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-alephalpha/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-alephalpha/tests/test_safety_hooks.py new file mode 100644 index 0000000000..d7df0126a3 --- /dev/null +++ b/packages/opentelemetry-instrumentation-alephalpha/tests/test_safety_hooks.py @@ -0,0 +1,143 @@ +from types import SimpleNamespace + +import pytest +from aleph_alpha_client import Prompt + +from opentelemetry.instrumentation.alephalpha.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_request_prompt(): + exporter, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[PII:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + + request = SimpleNamespace( + prompt=Prompt.from_json( + [ + {"type": "text", "data": "secret-a"}, + {"type": "text", "data": "secret-b"}, + {"type": "token_ids", "data": [1, 2, 3]}, + ] + ) + ) + with tracer.start_as_current_span("alephalpha.completion") as span: + _apply_prompt_safety(span, (request,), {}, "alephalpha.completion") + + prompt_json = request.prompt.to_json() + assert prompt_json[0]["data"] == "[PII:secret-a]" + assert prompt_json[1]["data"] == "[PII:secret-b]" + assert prompt_json[2]["data"] == [1, 2, 3] + assert len(exporter.get_finished_spans()[0].events) == 2 + + +def test_completion_safety_masks_completion_text(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[SECRET:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + response = SimpleNamespace( + completions=[ + SimpleNamespace(completion="secret-a"), + SimpleNamespace(completion="secret-b"), + ] + ) + with tracer.start_as_current_span("alephalpha.completion") as span: + _apply_completion_safety(span, response, "alephalpha.completion") + + assert response.completions[0].completion == "[SECRET:secret-a]" + assert response.completions[1].completion == "[SECRET:secret-b]" + assert len(exporter.get_finished_spans()[0].events) == 2 + + +def test_helpers_cover_passthrough_and_prompt_extraction_branches(): + _, tracer = _test_span() + + class BrokenRequest: + @property + def prompt(self): + raise RuntimeError("boom") + + class BrokenResponse: + @property + def completions(self): + raise RuntimeError("boom") + + request = SimpleNamespace(prompt=Prompt.from_text("plain")) + broken_request = BrokenRequest() + with tracer.start_as_current_span("alephalpha.completion") as span: + assert _apply_prompt_safety(span, (request,), {}, "alephalpha.completion") == ( + (request,), + {}, + ) + _apply_completion_safety( + span, + SimpleNamespace(completions=[]), + "alephalpha.completion", + ) + _apply_completion_safety( + span, + SimpleNamespace(completions=[SimpleNamespace(completion=None)]), + "alephalpha.completion", + ) + assert _apply_prompt_safety( + span, (broken_request,), {}, "alephalpha.completion" + ) == ((broken_request,), {}) + _apply_completion_safety(span, BrokenResponse(), "alephalpha.completion") + dict_request = SimpleNamespace(prompt={"type": "text", "data": "plain"}) + assert _apply_prompt_safety( + span, + (dict_request,), + {}, + "alephalpha.completion", + ) == ((dict_request,), {}) + + assert _resolve_masked_text("same", None) == ("same", False) + unchanged = SafetyResult(text="same", overall_action="MASK", findings=[]) + assert _resolve_masked_text("same", unchanged) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/__init__.py b/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/__init__.py index 885dd92150..9f57d3c628 100644 --- a/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/__init__.py +++ b/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/__init__.py @@ -26,6 +26,12 @@ from opentelemetry.instrumentation.bedrock.reusable_streaming_body import ( ReusableStreamingBody, ) +from opentelemetry.instrumentation.bedrock.safety import ( + _apply_converse_completion_safety, + _apply_converse_prompt_safety, + _apply_invoke_prompt_safety, + _prepare_invoke_response, +) from opentelemetry.instrumentation.bedrock.span_utils import ( converse_usage_record, set_converse_input_prompt_span_attributes, @@ -203,6 +209,7 @@ def with_instrumentation(*args, **kwargs): with tracer.start_as_current_span( _BEDROCK_INVOKE_SPAN_NAME, kind=SpanKind.CLIENT ) as span: + kwargs = _apply_invoke_prompt_safety(span, kwargs, _BEDROCK_INVOKE_SPAN_NAME) response = fn(*args, **kwargs) _handle_call(span, kwargs, response, metric_params, event_logger) return response @@ -240,6 +247,7 @@ def with_instrumentation(*args, **kwargs): with tracer.start_as_current_span( _BEDROCK_CONVERSE_SPAN_NAME, kind=SpanKind.CLIENT ) as span: + kwargs = _apply_converse_prompt_safety(span, kwargs, _BEDROCK_CONVERSE_SPAN_NAME) response = fn(*args, **kwargs) _handle_converse(span, kwargs, response, metric_params, event_logger) @@ -316,7 +324,7 @@ def _handle_call(span: Span, kwargs, response, metric_params, event_logger): response["body"]._raw_stream, response["body"]._content_length ) request_body = json.loads(kwargs.get("body")) - response_body = json.loads(response.get("body").read()) + response_body = _prepare_invoke_response(span, response, _BEDROCK_INVOKE_SPAN_NAME) headers = {} if "ResponseMetadata" in response: headers = response.get("ResponseMetadata").get("HTTPHeaders", {}) @@ -352,6 +360,7 @@ def _handle_call(span: Span, kwargs, response, metric_params, event_logger): @dont_throw def _handle_converse(span, kwargs, response, metric_params, event_logger): + _apply_converse_completion_safety(span, response, _BEDROCK_CONVERSE_SPAN_NAME) (provider, model_vendor, model) = _get_vendor_model(kwargs.get("modelId")) guardrail_converse(span, response, provider, model, metric_params) diff --git a/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/safety.py b/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/safety.py new file mode 100644 index 0000000000..2c894fb438 --- /dev/null +++ b/packages/opentelemetry-instrumentation-bedrock/opentelemetry/instrumentation/bedrock/safety.py @@ -0,0 +1,434 @@ +from __future__ import annotations + +import json +from io import BytesIO + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + run_completion_safety, + run_prompt_safety, +) +from opentelemetry.instrumentation.bedrock.reusable_streaming_body import ( + ReusableStreamingBody, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "Bedrock" + + +def _apply_invoke_prompt_safety(span, kwargs, span_name): + payload, as_bytes = _decode_payload(kwargs.get("body")) + if payload is None: + return kwargs + + masked_payload, changed = _mask_prompt_payload( + span, + payload, + span_name=span_name, + request_type=_request_type(span_name), + segment_index=0, + ) + if not changed: + return kwargs + + mutated_kwargs = dict(kwargs) + mutated_kwargs["body"] = _encode_payload(masked_payload, as_bytes) + return mutated_kwargs + + +def _apply_converse_prompt_safety(span, kwargs, span_name): + try: + mutated_kwargs = kwargs + + system_messages = kwargs.get("system") + if isinstance(system_messages, list): + mutated_system = None + for index, message in enumerate(system_messages): + text = message.get("text") if isinstance(message, dict) else None + if not isinstance(text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + text, + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=index, + segment_role="system", + ) + if not changed: + continue + if mutated_system is None: + mutated_kwargs = dict(kwargs) + mutated_system = [dict(item) for item in system_messages] + mutated_kwargs["system"] = mutated_system + mutated_system[index]["text"] = updated_text + + messages = kwargs.get("messages") + if not isinstance(messages, list): + return mutated_kwargs + + mutated_messages = None + for index, message in enumerate(messages): + if not isinstance(message, dict): + continue + updated_content, changed = _mask_prompt_content( + span, + message.get("content"), + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=index, + segment_role=message.get("role") or "user", + ) + if not changed: + continue + if mutated_messages is None: + if mutated_kwargs is kwargs: + mutated_kwargs = dict(kwargs) + mutated_messages = [dict(item) if isinstance(item, dict) else item for item in messages] + mutated_kwargs["messages"] = mutated_messages + mutated_messages[index]["content"] = updated_content + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_invoke_completion_safety(span, raw_response, span_name): + payload, as_bytes = _decode_payload(raw_response) + if payload is None: + return raw_response, False + + masked_payload, changed = _mask_completion_payload( + span, + payload, + span_name=span_name, + request_type=_request_type(span_name), + segment_index=0, + ) + if not changed: + return raw_response, False + return _encode_payload(masked_payload, as_bytes), True + + +def _prepare_invoke_response(span, response, span_name): + body = response.get("body") + if body is None: + return None + + response["body"] = ReusableStreamingBody(body._raw_stream, body._content_length) + raw_response = response["body"].read() + masked_response, changed = _apply_invoke_completion_safety( + span, raw_response, span_name + ) + if changed: + raw_response = masked_response + response["body"] = ReusableStreamingBody(BytesIO(raw_response), len(raw_response)) + return json.loads(raw_response) + + +def _apply_converse_completion_safety(span, response, span_name): + try: + output = response.get("output") + if not isinstance(output, dict): + return + message = output.get("message") + if not isinstance(message, dict): + return + updated_content, changed = _mask_completion_content( + span, + message.get("content"), + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=0, + ) + if changed: + updated_message = dict(message) + updated_message["content"] = updated_content + updated_output = dict(output) + updated_output["message"] = updated_message + response["output"] = updated_output + except Exception: + return + + +def _mask_prompt_payload(span, value, *, span_name, request_type, segment_index): + if isinstance(value, dict): + updated = value + for key, item in value.items(): + if key in {"prompt", "inputText", "text"} and isinstance(item, str): + updated_item, changed = _mask_prompt_text( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role="user", + ) + elif key in {"messages", "content"}: + updated_item, changed = _mask_prompt_content( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role="user", + ) + else: + updated_item, changed = _mask_prompt_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + if not changed: + continue + if updated is value: + updated = dict(value) + updated[key] = updated_item + return updated, updated is not value + + if isinstance(value, list): + updated = value + for index, item in enumerate(value): + updated_item, changed = _mask_prompt_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if not changed: + continue + if updated is value: + updated = list(value) + updated[index] = updated_item + return updated, updated is not value + + return value, False + + +def _mask_completion_payload(span, value, *, span_name, request_type, segment_index): + if isinstance(value, dict): + updated = value + for key, item in value.items(): + if key in {"completion", "outputText", "generated_text", "text"} and isinstance(item, str): + updated_item, changed = _mask_completion_text( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + elif key in {"content", "completions", "generations"}: + updated_item, changed = _mask_completion_content( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + else: + updated_item, changed = _mask_completion_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + if not changed: + continue + if updated is value: + updated = dict(value) + updated[key] = updated_item + return updated, updated is not value + + if isinstance(value, list): + updated = value + for index, item in enumerate(value): + updated_item, changed = _mask_completion_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if not changed: + continue + if updated is value: + updated = list(value) + updated[index] = updated_item + return updated, updated is not value + + return value, False + + +def _mask_prompt_content( + span, + content, + *, + span_name, + request_type, + segment_index, + segment_role, +): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated = content + for index, item in enumerate(content): + if isinstance(item, dict) and isinstance(item.get("text"), str): + updated_text, changed = _mask_prompt_text( + span, + item["text"], + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated is content: + updated = [dict(block) if isinstance(block, dict) else block for block in content] + updated[index]["text"] = updated_text + continue + + updated_item, changed = _mask_prompt_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if not changed: + continue + if updated is content: + updated = list(content) + updated[index] = updated_item + + return updated, updated is not content + + +def _mask_completion_content(span, content, *, span_name, request_type, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated = content + for index, item in enumerate(content): + if isinstance(item, dict) and isinstance(item.get("text"), str): + updated_text, changed = _mask_completion_text( + span, + item["text"], + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + ) + if not changed: + continue + if updated is content: + updated = [dict(block) if isinstance(block, dict) else block for block in content] + updated[index]["text"] = updated_text + continue + + updated_item, changed = _mask_completion_payload( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if not changed: + continue + if updated is content: + updated = list(content) + updated[index] = updated_item + + return updated, updated is not content + + +def _mask_prompt_text( + span, + text, + *, + span_name, + request_type, + segment_index, + segment_role, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, request_type, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _decode_payload(raw_value): + try: + if isinstance(raw_value, bytes): + return json.loads(raw_value.decode("utf-8")), True + if isinstance(raw_value, str): + return json.loads(raw_value), False + except Exception: + return None, False + return None, False + + +def _encode_payload(payload, as_bytes): + encoded = json.dumps(payload).encode("utf-8") + if as_bytes: + return encoded + return encoded.decode("utf-8") + + +def _request_type(span_name): + if "converse" in span_name: + return LLMRequestTypeValues.CHAT.value + return LLMRequestTypeValues.COMPLETION.value + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-bedrock/pyproject.toml b/packages/opentelemetry-instrumentation-bedrock/pyproject.toml index be083e511a..b779546fe5 100644 --- a/packages/opentelemetry-instrumentation-bedrock/pyproject.toml +++ b/packages/opentelemetry-instrumentation-bedrock/pyproject.toml @@ -13,6 +13,7 @@ requires-python = ">=3.10,<4" dependencies = [ "anthropic>=0.17.0", "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -70,5 +71,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-bedrock/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-bedrock/tests/test_safety_hooks.py new file mode 100644 index 0000000000..43edaf77ba --- /dev/null +++ b/packages/opentelemetry-instrumentation-bedrock/tests/test_safety_hooks.py @@ -0,0 +1,224 @@ +import json +from io import BytesIO + +import pytest + +from opentelemetry.instrumentation.bedrock.reusable_streaming_body import ( + ReusableStreamingBody, +) +from opentelemetry.instrumentation.bedrock.safety import ( + _apply_converse_completion_safety, + _apply_converse_prompt_safety, + _decode_payload, + _encode_payload, + _apply_invoke_completion_safety, + _apply_invoke_prompt_safety, + _mask_completion_payload, + _mask_prompt_payload, + _prepare_invoke_response, + _request_type, + _resolve_masked_text, +) +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_invoke_prompt_safety_masks_json_body(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"body": json.dumps({"prompt": "secret"}), "modelId": "anthropic.claude"} + with tracer.start_as_current_span("bedrock.completion") as span: + updated_kwargs = _apply_invoke_prompt_safety(span, kwargs, "bedrock.completion") + + assert json.loads(updated_kwargs["body"])["prompt"] == "[PII.prompt]" + + +def test_converse_prompt_safety_masks_message_content(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"messages": [{"role": "user", "content": [{"text": "secret"}]}]} + with tracer.start_as_current_span("bedrock.converse") as span: + updated_kwargs = _apply_converse_prompt_safety(span, kwargs, "bedrock.converse") + + assert updated_kwargs["messages"][0]["content"][0]["text"] == "[PII.chat]" + + +def test_invoke_completion_safety_masks_json_response(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + raw_response = json.dumps({"completion": "secret"}) + with tracer.start_as_current_span("bedrock.completion") as span: + updated_response, changed = _apply_invoke_completion_safety( + span, raw_response, "bedrock.completion" + ) + + assert changed is True + assert json.loads(updated_response)["completion"] == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prepare_invoke_response_masks_and_rebuilds_streaming_body(): + _, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + raw_response = json.dumps({"completion": "secret"}).encode("utf-8") + response = { + "body": ReusableStreamingBody(BytesIO(raw_response), len(raw_response)), + } + + with tracer.start_as_current_span("bedrock.completion") as span: + parsed_response = _prepare_invoke_response(span, response, "bedrock.completion") + + assert parsed_response["completion"] == "[SECRET.output]" + assert json.loads(response["body"].read().decode("utf-8"))["completion"] == "[SECRET.output]" + + +def test_converse_helpers_cover_system_messages_and_completion_output(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + kwargs = { + "system": [{"text": "system-secret"}, {"ignore": "x"}], + "messages": [{"role": "user", "content": [{"text": "message-secret"}, {"image": "ignored"}]}], + } + response = {"output": {"message": {"content": [{"text": "completion-secret"}]}}} + with tracer.start_as_current_span("bedrock.converse") as span: + updated_kwargs = _apply_converse_prompt_safety(span, kwargs, "bedrock.converse") + _apply_converse_completion_safety(span, response, "bedrock.converse") + + assert updated_kwargs["system"][0]["text"] == "[MASKED:system-secret]" + assert updated_kwargs["messages"][0]["content"][0]["text"] == "[MASKED:message-secret]" + assert response["output"]["message"]["content"][0]["text"] == "[MASKED:completion-secret]" + + +def test_payload_helpers_cover_nested_structures_and_invalid_inputs(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + prompt_payload = {"items": [{"prompt": "prompt-secret"}], "messages": [{"content": [{"text": "message-secret"}]}]} + completion_payload = {"items": [{"completion": "completion-secret"}], "generations": [[{"text": "gen-secret"}]]} + with tracer.start_as_current_span("bedrock.completion") as span: + masked_prompt, prompt_changed = _mask_prompt_payload( + span, + prompt_payload, + span_name="bedrock.completion", + request_type="completion", + segment_index=0, + ) + masked_completion, completion_changed = _mask_completion_payload( + span, + completion_payload, + span_name="bedrock.completion", + request_type="completion", + segment_index=0, + ) + + assert prompt_changed is True + assert masked_prompt["items"][0]["prompt"] == "[MASKED:prompt-secret]" + assert masked_prompt["messages"][0]["content"][0]["text"] == "[MASKED:message-secret]" + assert completion_changed is True + assert masked_completion["items"][0]["completion"] == "[MASKED:completion-secret]" + assert masked_completion["generations"][0][0]["text"] == "[MASKED:gen-secret]" + assert _decode_payload("not-json") == (None, False) + assert _decode_payload(123) == (None, False) + assert _encode_payload({"x": 1}, False) == '{"x": 1}' + assert _encode_payload({"x": 1}, True) == b'{"x": 1}' + assert _request_type("bedrock.converse") == "chat" + assert _request_type("bedrock.completion") == "completion" + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/__init__.py b/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/__init__.py index daff471dfe..54e465488d 100644 --- a/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/__init__.py +++ b/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/__init__.py @@ -13,6 +13,10 @@ emit_message_events, emit_streaming_response_events, ) +from opentelemetry.instrumentation.groq.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.groq.span_utils import ( set_input_attributes, set_model_input_attributes, @@ -250,6 +254,7 @@ def _wrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, name) _handle_input(span, kwargs, event_logger) start_time = time.time() @@ -289,6 +294,7 @@ def _wrap( attributes=metric_attributes, ) + _apply_completion_safety(span, response, name) _handle_response(span, response, token_histogram, event_logger) except Exception as ex: # pylint: disable=broad-except @@ -332,6 +338,7 @@ async def _awrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, name) _handle_input(span, kwargs, event_logger) start_time = time.time() @@ -371,6 +378,7 @@ async def _awrap( attributes=metric_attributes, ) + _apply_completion_safety(span, response, name) _handle_response(span, response, token_histogram, event_logger) if span.is_recording(): diff --git a/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/safety.py b/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/safety.py new file mode 100644 index 0000000000..43007f1699 --- /dev/null +++ b/packages/opentelemetry-instrumentation-groq/opentelemetry/instrumentation/groq/safety.py @@ -0,0 +1,205 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "Groq" + + +def _apply_prompt_safety(span, kwargs, span_name): + try: + mutated_kwargs = kwargs + + prompt = kwargs.get("prompt") + if isinstance(prompt, str): + updated_prompt, changed = _mask_prompt_text( + span, + prompt, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + mutated_kwargs = dict(kwargs) + mutated_kwargs["prompt"] = updated_prompt + + messages = kwargs.get("messages") + if not isinstance(messages, list): + return mutated_kwargs + + mutated_messages = None + for index, message in enumerate(messages): + updated_content, changed = _mask_prompt_content( + span, + get_object_value(message, "content"), + span_name=span_name, + segment_index=index, + segment_role=get_object_value(message, "role"), + ) + if not changed: + continue + if mutated_messages is None: + if mutated_kwargs is kwargs: + mutated_kwargs = dict(kwargs) + mutated_messages = clone_value(messages) + mutated_kwargs["messages"] = mutated_messages + set_object_value(mutated_messages[index], "content", updated_content) + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_completion_safety(span, response, span_name): + try: + choices = get_object_value(response, "choices") or [] + for index, choice in enumerate(choices): + message = get_object_value(choice, "message") + if message is not None: + updated_content, changed = _mask_completion_content( + span, + get_object_value(message, "content"), + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(message, "content", updated_content) + continue + + text = get_object_value(choice, "text") + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + span, + text, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(choice, "text", updated_text) + except Exception: + return + + +def _mask_prompt_content(span, content, *, span_name, segment_index, segment_role): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text", "input_text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content(span, content, *, span_name, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text", "output_text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + span_name=span_name, + segment_index=segment_index, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + span_name, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index, metadata=None): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=segment_index, + segment_role="assistant", + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-groq/pyproject.toml b/packages/opentelemetry-instrumentation-groq/pyproject.toml index e1182515d9..18c9de989f 100644 --- a/packages/opentelemetry-instrumentation-groq/pyproject.toml +++ b/packages/opentelemetry-instrumentation-groq/pyproject.toml @@ -11,6 +11,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -71,5 +72,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-groq/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-groq/tests/test_safety_hooks.py new file mode 100644 index 0000000000..e7a818f69a --- /dev/null +++ b/packages/opentelemetry-instrumentation-groq/tests/test_safety_hooks.py @@ -0,0 +1,189 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.groq import _awrap +from opentelemetry.instrumentation.groq.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _mask_completion_content, + _mask_prompt_content, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv_ai import SpanAttributes + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_chat_message(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"messages": [{"role": "user", "content": "secret"}]} + with tracer.start_as_current_span("groq.chat") as span: + updated_kwargs = _apply_prompt_safety(span, kwargs, "groq.chat") + + assert updated_kwargs["messages"][0]["content"] == "[PII.chat]" + + +def test_completion_safety_masks_chat_choice(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.chat]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="secret"))]) + with tracer.start_as_current_span("groq.chat") as span: + _apply_completion_safety(span, response, "groq.chat") + + assert response.choices[0].message.content == "[SECRET.chat]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_safety_masks_prompt_and_block_content(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + + kwargs = { + "prompt": "secret-prompt", + "messages": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "secret-block"}, + {"type": "image_url", "text": "ignored"}, + ], + } + ], + } + with tracer.start_as_current_span("groq.chat") as span: + updated_kwargs = _apply_prompt_safety(span, kwargs, "groq.chat") + + assert updated_kwargs["prompt"] == "[MASKED:secret-prompt]" + assert updated_kwargs["messages"][0]["content"][0]["text"] == "[MASKED:secret-block]" + assert kwargs["messages"][0]["content"][0]["text"] == "secret-block" + + +def test_completion_helpers_cover_text_and_passthrough_paths(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + response = SimpleNamespace( + choices=[ + SimpleNamespace(text="secret-text"), + SimpleNamespace(message=SimpleNamespace(content=[{"type": "output_text", "text": "secret-block"}])), + ] + ) + with tracer.start_as_current_span("groq.chat") as span: + _apply_completion_safety(span, response, "groq.chat") + assert _mask_prompt_content(span, 123, span_name="groq.chat", segment_index=0, segment_role="user") == ( + 123, + False, + ) + assert _mask_completion_content(span, None, span_name="groq.chat", segment_index=0) == ( + None, + False, + ) + + assert response.choices[0].text == "[MASKED:secret-text]" + assert response.choices[1].message.content[0]["text"] == "[MASKED:secret-block]" + assert _resolve_masked_text("same", None) == ("same", False) + unchanged = SafetyResult(text="same", overall_action="MASK", findings=[]) + assert _resolve_masked_text("same", unchanged) == ("same", False) + assert len(exporter.get_finished_spans()[0].events) >= 2 + + +@pytest.mark.asyncio +async def test_async_wrapper_masks_completion_without_metrics_histogram(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[MASKED:secret]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + async def wrapped(*args, **kwargs): + return { + "model": "groq-test", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "secret"}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + + wrapper = _awrap(tracer, None, None, None, None, {"span_name": "groq.chat"}) + response = await wrapper( + wrapped, + None, + (), + {"messages": [{"role": "user", "content": "prompt"}]}, + ) + + assert response["choices"][0]["message"]["content"] == "[MASKED:secret]" + span = exporter.get_finished_spans()[0] + assert span.attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[MASKED:secret]" diff --git a/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/__init__.py b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/__init__.py index 1eddfbcb3d..1a2defc272 100644 --- a/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/__init__.py +++ b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/__init__.py @@ -12,6 +12,10 @@ TraceloopCallbackHandler, ) from opentelemetry.instrumentation.langchain.config import Config +from opentelemetry.instrumentation.langchain.safety import ( + instrument_safety_wrappers, + uninstrument_safety_wrappers, +) from opentelemetry.instrumentation.langchain.utils import is_package_available from opentelemetry.instrumentation.langchain.version import __version__ from opentelemetry.instrumentation.utils import unwrap @@ -95,6 +99,7 @@ def _instrument(self, **kwargs): name="BaseCallbackManager.__init__", wrapper=_BaseCallbackManagerInitWrapper(traceloopCallbackHandler), ) + instrument_safety_wrappers() if not self.disable_trace_context_propagation: self._wrap_openai_functions_for_tracing(traceloopCallbackHandler) @@ -181,6 +186,7 @@ def _wrap_openai_functions_for_tracing(self, traceloopCallbackHandler): def _uninstrument(self, **kwargs): unwrap("langchain_core.callbacks", "BaseCallbackManager.__init__") + uninstrument_safety_wrappers() if not self.disable_trace_context_propagation: if is_package_available("langchain_community"): unwrap("langchain_community.llms.openai", "BaseOpenAI._generate") diff --git a/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py new file mode 100644 index 0000000000..66b0ad0450 --- /dev/null +++ b/packages/opentelemetry-instrumentation-langchain/opentelemetry/instrumentation/langchain/safety.py @@ -0,0 +1,488 @@ +from __future__ import annotations + +from typing import Any + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.instrumentation.langchain.vendor_detection import ( + detect_vendor_from_class, +) +from opentelemetry.instrumentation.utils import unwrap +from opentelemetry.semconv_ai import LLMRequestTypeValues +from wrapt import wrap_function_wrapper + +PROVIDER = "Langchain" + + +def base_chat_model_generate_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_chat_prompt_safety(instance, args, kwargs) + return wrapped(*updated_args, **updated_kwargs) + + +async def base_chat_model_agenerate_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_chat_prompt_safety(instance, args, kwargs) + return await wrapped(*updated_args, **updated_kwargs) + + +def base_llm_generate_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_llm_prompt_safety(instance, args, kwargs) + return wrapped(*updated_args, **updated_kwargs) + + +async def base_llm_agenerate_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_llm_prompt_safety(instance, args, kwargs) + return await wrapped(*updated_args, **updated_kwargs) + + +def base_chat_model_generate_with_cache_wrapper(wrapped, instance, args, kwargs): + response = wrapped(*args, **kwargs) + _apply_chat_result_completion_safety(instance, response) + return response + + +async def base_chat_model_agenerate_with_cache_wrapper( + wrapped, instance, args, kwargs +): + response = await wrapped(*args, **kwargs) + _apply_chat_result_completion_safety(instance, response) + return response + + +def base_llm_generate_helper_wrapper(wrapped, instance, args, kwargs): + response = wrapped(*args, **kwargs) + _apply_llm_result_completion_safety(instance, response) + return response + + +async def base_llm_agenerate_helper_wrapper(wrapped, instance, args, kwargs): + response = await wrapped(*args, **kwargs) + _apply_llm_result_completion_safety(instance, response) + return response + + +def _apply_chat_prompt_safety(instance, args, kwargs): + messages = args[0] if args else kwargs.get("messages") + if not isinstance(messages, list): + return args, kwargs + + provider = _provider_name(instance) + span_name = f"{instance.__class__.__name__}.chat" + updated_batches = messages + changed = False + + for batch_index, batch in enumerate(messages): + if not isinstance(batch, list): + continue + updated_batch = batch + for message_index, message in enumerate(batch): + content = get_object_value(message, "content") + updated_content, content_changed = _mask_prompt_content( + content, + provider=provider, + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=message_index, + segment_role=_message_role(message), + metadata={"batch_index": batch_index}, + ) + if not content_changed: + continue + if updated_batches is messages: + updated_batches = clone_value(messages) + if updated_batch is batch: + updated_batch = updated_batches[batch_index] + set_object_value(updated_batch[message_index], "content", updated_content) + changed = True + + if not changed: + return args, kwargs + + if args: + updated_args = list(args) + updated_args[0] = updated_batches + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["messages"] = updated_batches + return args, updated_kwargs + + +def _apply_llm_prompt_safety(instance, args, kwargs): + prompts = args[0] if args else kwargs.get("prompts") + if not isinstance(prompts, list): + return args, kwargs + + provider = _provider_name(instance) + span_name = f"{instance.__class__.__name__}.completion" + updated_prompts = prompts + changed = False + + for index, prompt in enumerate(prompts): + if not isinstance(prompt, str): + continue + updated_prompt, prompt_changed = _mask_prompt_text( + prompt, + provider=provider, + span_name=span_name, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=index, + segment_role="user", + ) + if not prompt_changed: + continue + if updated_prompts is prompts: + updated_prompts = list(prompts) + updated_prompts[index] = updated_prompt + changed = True + + if not changed: + return args, kwargs + + if args: + updated_args = list(args) + updated_args[0] = updated_prompts + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["prompts"] = updated_prompts + return args, updated_kwargs + + +def _apply_chat_result_completion_safety(instance, response): + generations = get_object_value(response, "generations") + if not isinstance(generations, list): + return + + provider = _provider_name(instance) + span_name = f"{instance.__class__.__name__}.chat" + for index, generation in enumerate(generations): + message = get_object_value(generation, "message") + message_content = get_object_value(message, "content") if message is not None else None + if message is not None: + updated_content, changed = _mask_completion_content( + message_content, + provider=provider, + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=index, + segment_role="assistant", + ) + if changed: + set_object_value(message, "content", updated_content) + if isinstance(updated_content, str): + set_object_value(generation, "text", updated_content) + continue + text = get_object_value(generation, "text") + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + text, + provider=provider, + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=index, + segment_role="assistant", + ) + if changed: + set_object_value(generation, "text", updated_text) + if message is not None and isinstance(message_content, str): + set_object_value(message, "content", updated_text) + + +def _apply_llm_result_completion_safety(instance, response): + generations = get_object_value(response, "generations") + if not isinstance(generations, list): + return + + provider = _provider_name(instance) + span_name = f"{instance.__class__.__name__}.completion" + for batch in generations: + if not isinstance(batch, list): + continue + for index, generation in enumerate(batch): + text = get_object_value(generation, "text") + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + text, + provider=provider, + span_name=span_name, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=index, + segment_role="assistant", + ) + if changed: + set_object_value(generation, "text", updated_text) + message = get_object_value(generation, "message") + if message is not None and isinstance(get_object_value(message, "content"), str): + set_object_value(message, "content", updated_text) + + +def _mask_prompt_content( + content, + *, + provider, + span_name, + request_type, + segment_index, + segment_role, + metadata=None, +): + if isinstance(content, str): + return _mask_prompt_text( + content, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_prompt_text( + block, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + text = _content_text(block) + if not isinstance(text, str): + continue + updated_text, changed = _mask_prompt_text( + text, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + _set_content_text(updated_content[block_index], updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content( + content, + *, + provider, + span_name, + request_type, + segment_index, + segment_role, +): + if isinstance(content, str): + return _mask_completion_text( + content, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_completion_text( + block, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + text = _content_text(block) + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + text, + provider=provider, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + _set_content_text(updated_content[block_index], updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + text, + *, + provider, + span_name, + request_type, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=None, + provider=provider, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text( + text, + *, + provider, + span_name, + request_type, + segment_index, + segment_role, +): + result = run_completion_safety( + span=None, + provider=provider, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _provider_name(instance) -> str: + return detect_vendor_from_class(instance.__class__.__name__) or PROVIDER + + +def _message_role(message) -> str: + message_type = get_object_value(message, "type") + if message_type == "human": + return "user" + if message_type == "system": + return "system" + if message_type == "ai": + return "assistant" + if message_type == "tool": + return "tool" + return str(get_object_value(message, "role", "unknown")).lower() + + +def _content_text(block): + if get_object_value(block, "type") == "text": + return get_object_value(block, "text") + if get_object_value(block, "text") is not None: + return get_object_value(block, "text") + return None + + +def _set_content_text(block, value): + if get_object_value(block, "type") == "text": + return set_object_value(block, "text", value) + return set_object_value(block, "text", value) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True + + +_WRAPPED_METHODS = ( + ( + "langchain_core.language_models.chat_models", + "BaseChatModel.generate", + base_chat_model_generate_wrapper, + ), + ( + "langchain_core.language_models.chat_models", + "BaseChatModel.agenerate", + base_chat_model_agenerate_wrapper, + ), + ( + "langchain_core.language_models.chat_models", + "BaseChatModel._generate_with_cache", + base_chat_model_generate_with_cache_wrapper, + ), + ( + "langchain_core.language_models.chat_models", + "BaseChatModel._agenerate_with_cache", + base_chat_model_agenerate_with_cache_wrapper, + ), + ( + "langchain_core.language_models.llms", + "BaseLLM.generate", + base_llm_generate_wrapper, + ), + ( + "langchain_core.language_models.llms", + "BaseLLM.agenerate", + base_llm_agenerate_wrapper, + ), + ( + "langchain_core.language_models.llms", + "BaseLLM._generate_helper", + base_llm_generate_helper_wrapper, + ), + ( + "langchain_core.language_models.llms", + "BaseLLM._agenerate_helper", + base_llm_agenerate_helper_wrapper, + ), +) + + +def instrument_safety_wrappers(): + for module_name, function_name, wrapper in _WRAPPED_METHODS: + wrap_function_wrapper(module_name, function_name, wrapper) + + +def uninstrument_safety_wrappers(): + for module_name, function_name, _ in _WRAPPED_METHODS: + unwrap(module_name, function_name) diff --git a/packages/opentelemetry-instrumentation-langchain/pyproject.toml b/packages/opentelemetry-instrumentation-langchain/pyproject.toml index 1a7600fcfb..0da2ccd8f4 100644 --- a/packages/opentelemetry-instrumentation-langchain/pyproject.toml +++ b/packages/opentelemetry-instrumentation-langchain/pyproject.toml @@ -12,6 +12,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -89,3 +90,9 @@ select = ["E", "F", "W"] [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "langgraph-checkpoint>=4.0.0", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } + +[tool.pytest.ini_options] +markers = ["safety: tests for FortifyRoot safety masking hooks"] diff --git a/packages/opentelemetry-instrumentation-langchain/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-langchain/tests/test_safety_hooks.py new file mode 100644 index 0000000000..1e54c04731 --- /dev/null +++ b/packages/opentelemetry-instrumentation-langchain/tests/test_safety_hooks.py @@ -0,0 +1,342 @@ +from langchain_core.messages import AIMessage, HumanMessage +from langchain_core.outputs import ChatGeneration, ChatResult, Generation, LLMResult +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.langchain.safety import ( + _apply_chat_prompt_safety, + _apply_chat_result_completion_safety, + _apply_llm_prompt_safety, + _apply_llm_result_completion_safety, + _content_text, + _message_role, + _set_content_text, + base_chat_model_agenerate_with_cache_wrapper, + base_chat_model_agenerate_wrapper, + base_chat_model_generate_with_cache_wrapper, + base_chat_model_generate_wrapper, + base_llm_agenerate_helper_wrapper, + base_llm_agenerate_wrapper, + base_llm_generate_helper_wrapper, + base_llm_generate_wrapper, + instrument_safety_wrappers, + uninstrument_safety_wrappers, + _provider_name, + _resolve_masked_text, +) + +pytestmark = pytest.mark.safety + + +class _FakeChatModel: + pass + + +class _FakeLLM: + pass + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def test_chat_prompt_safety_masks_messages_before_wrapped_call(): + captured = {} + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.langchain]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="PII", + severity="MEDIUM", + action="MASK", + rule_name="PII.langchain", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + def wrapped(messages, **kwargs): + captured["messages"] = messages + return "ok" + + messages = [[HumanMessage(content="secret")]] + result = base_chat_model_generate_wrapper( + wrapped, + _FakeChatModel(), + (messages,), + {}, + ) + + assert result == "ok" + assert captured["messages"][0][0].content == "[PII.langchain]" + assert messages[0][0].content == "secret" + + +def test_llm_prompt_safety_masks_prompts_before_wrapped_call(): + captured = {} + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.langchain]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="PII", + severity="MEDIUM", + action="MASK", + rule_name="PII.langchain", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + def wrapped(prompts, *args, **kwargs): + captured["prompts"] = prompts + return "ok" + + result = base_llm_generate_wrapper( + wrapped, + _FakeLLM(), + (["secret"], None, None), + {}, + ) + + assert result == "ok" + assert captured["prompts"] == ["[PII.langchain]"] + + +def test_chat_completion_safety_masks_chat_result_before_callbacks(): + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.langchain]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="SECRET", + severity="HIGH", + action="MASK", + rule_name="SECRET.langchain", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + def wrapped(*args, **kwargs): + return ChatResult(generations=[ChatGeneration(message=AIMessage(content="secret"))]) + + response = base_chat_model_generate_with_cache_wrapper( + wrapped, + _FakeChatModel(), + ([],), + {}, + ) + + assert response.generations[0].message.content == "[SECRET.langchain]" + assert response.generations[0].text == "[SECRET.langchain]" + + +def test_llm_completion_safety_masks_llm_result_before_callbacks(): + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.langchain]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="SECRET", + severity="HIGH", + action="MASK", + rule_name="SECRET.langchain", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + def wrapped(*args, **kwargs): + return LLMResult(generations=[[Generation(text="secret")]]) + + response = base_llm_generate_helper_wrapper( + wrapped, + _FakeLLM(), + (["prompt"], None, []), + {}, + ) + + assert response.generations[0][0].text == "[SECRET.langchain]" + + +@pytest.mark.asyncio +async def test_async_wrappers_cover_kwargs_and_cache_paths(): + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "MEDIUM", "MASK", "PII.langchain", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.langchain", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + async def wrapped_chat(messages, **kwargs): + return messages + + async def wrapped_cache(*args, **kwargs): + return ChatResult( + generations=[ + ChatGeneration(message=AIMessage(content=[{"type": "text", "text": "secret-block"}])) + ] + ) + + async def wrapped_llm(prompts, *args, **kwargs): + return prompts + + async def wrapped_llm_cache(*args, **kwargs): + return LLMResult(generations=[[Generation(text="secret-llm")]]) + + chat_messages = [[HumanMessage(content=[{"type": "text", "text": "secret-block"}])]] + masked_messages = await base_chat_model_agenerate_wrapper( + wrapped_chat, + _FakeChatModel(), + (), + {"messages": chat_messages}, + ) + chat_response = await base_chat_model_agenerate_with_cache_wrapper( + wrapped_cache, + _FakeChatModel(), + ([],), + {}, + ) + masked_prompts = await base_llm_agenerate_wrapper( + wrapped_llm, + _FakeLLM(), + (), + {"prompts": ["secret-llm"]}, + ) + llm_response = await base_llm_agenerate_helper_wrapper( + wrapped_llm_cache, + _FakeLLM(), + (["prompt"], None, []), + {}, + ) + + assert masked_messages[0][0].content[0]["text"] == "[MASKED:secret-block]" + assert chat_response.generations[0].message.content[0]["text"] == "[MASKED:secret-block]" + assert chat_response.generations[0].text == "[MASKED:secret-block]" + assert masked_prompts == ["[MASKED:secret-llm]"] + assert llm_response.generations[0][0].text == "[MASKED:secret-llm]" + + +def test_registration_and_helper_fallbacks(): + with patch("opentelemetry.instrumentation.langchain.safety.wrap_function_wrapper") as wrap_mock, patch( + "opentelemetry.instrumentation.langchain.safety.unwrap" + ) as unwrap_mock: + instrument_safety_wrappers() + uninstrument_safety_wrappers() + + assert wrap_mock.call_count == 8 + assert unwrap_mock.call_count == 8 + assert _provider_name(_FakeChatModel()) == "Langchain" + assert _resolve_masked_text("same", None) == ("same", False) + assert _resolve_masked_text("same", SafetyResult(text="same", overall_action="MASK", findings=[])) == ( + "same", + False, + ) + + +def test_internal_helpers_cover_kwargs_passthrough_and_block_updates(): + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "MEDIUM", "MASK", "PII.langchain", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.langchain", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + chat_messages = [[HumanMessage(content=[{"type": "text", "text": "secret-block"}, "secret-inline"])]] + _, updated_kwargs = _apply_chat_prompt_safety(_FakeChatModel(), (), {"messages": chat_messages}) + assert updated_kwargs["messages"][0][0].content[0]["text"] == "[MASKED:secret-block]" + assert updated_kwargs["messages"][0][0].content[1] == "[MASKED:secret-inline]" + assert _apply_chat_prompt_safety(_FakeChatModel(), (), {"messages": "invalid"}) == ((), {"messages": "invalid"}) + + _, updated_llm_kwargs = _apply_llm_prompt_safety(_FakeLLM(), (), {"prompts": ["secret", 1]}) + assert updated_llm_kwargs["prompts"][0] == "[MASKED:secret]" + assert _apply_llm_prompt_safety(_FakeLLM(), (), {"prompts": "invalid"}) == ((), {"prompts": "invalid"}) + + chat_response = ChatResult( + generations=[ + ChatGeneration(message=AIMessage(content=[{"type": "text", "text": "secret-block"}])), + ChatGeneration(message=AIMessage(content="secret-text"), text="secret-text"), + ] + ) + _apply_chat_result_completion_safety(_FakeChatModel(), chat_response) + assert chat_response.generations[0].message.content[0]["text"] == "[MASKED:secret-block]" + assert chat_response.generations[0].text == "[MASKED:secret-block]" + assert chat_response.generations[1].text == "[MASKED:secret-text]" + + llm_response = SimpleNamespace( + generations=[ + [SimpleNamespace(text="secret-text", message=SimpleNamespace(content="secret-text"))], + "ignore", + ] + ) + _apply_llm_result_completion_safety(_FakeLLM(), llm_response) + assert llm_response.generations[0][0].text == "[MASKED:secret-text]" + assert llm_response.generations[0][0].message.content == "[MASKED:secret-text]" + + assert _message_role(SimpleNamespace(type="system")) == "system" + assert _message_role(SimpleNamespace(type="ai")) == "assistant" + assert _message_role(SimpleNamespace(type="tool")) == "tool" + block = {"type": "text", "text": "value"} + assert _content_text(block) == "value" + assert _set_content_text(block, "new") is True + assert block["text"] == "new" diff --git a/packages/opentelemetry-instrumentation-litellm/README.md b/packages/opentelemetry-instrumentation-litellm/README.md new file mode 100644 index 0000000000..1f91a0f744 --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/README.md @@ -0,0 +1,3 @@ +# OpenTelemetry LiteLLM instrumentation + +FortifyRoot safety-enabled instrumentation for LiteLLM non-stream completion calls. diff --git a/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/__init__.py b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/__init__.py new file mode 100644 index 0000000000..0f8f05cc6b --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/__init__.py @@ -0,0 +1,368 @@ +"""OpenTelemetry LiteLLM instrumentation.""" + +import inspect +import logging +from typing import Collection + +from opentelemetry import context as context_api +from opentelemetry.instrumentation.instrumentor import BaseInstrumentor +from opentelemetry.instrumentation.litellm.safety import ( + apply_completion_safety, + apply_prompt_safety, + extract_prompt_texts, + extract_text_content, +) +from opentelemetry.instrumentation.litellm.version import __version__ +from opentelemetry.instrumentation.utils import _SUPPRESS_INSTRUMENTATION_KEY, unwrap +from opentelemetry.semconv._incubating.attributes import ( + gen_ai_attributes as GenAIAttributes, +) +from opentelemetry.semconv_ai import ( + LLMRequestTypeValues, + SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY, + SpanAttributes, +) +from opentelemetry.trace import SpanKind, Status, StatusCode, get_tracer, use_span +from wrapt import wrap_function_wrapper + +logger = logging.getLogger(__name__) + +_instruments = ("litellm >= 1.71.2, < 2",) + +_WRAPPED_METHODS = [ + ("litellm", "completion", False, False), + ("litellm", "acompletion", True, False), + ("litellm", "text_completion", False, True), + ("litellm", "atext_completion", True, True), + ("litellm.main", "completion", False, False), + ("litellm.main", "acompletion", True, False), + ("litellm.main", "text_completion", False, True), + ("litellm.main", "atext_completion", True, True), +] + + +class LiteLLMInstrumentor(BaseInstrumentor): + def instrumentation_dependencies(self) -> Collection[str]: + return _instruments + + def _instrument(self, **kwargs): + tracer_provider = kwargs.get("tracer_provider") + tracer = get_tracer(__name__, __version__, tracer_provider) + + for module_name, func_name, is_async, is_text_completion in _WRAPPED_METHODS: + wrapper = ( + _build_async_wrapper(tracer, is_text_completion) + if is_async + else _build_sync_wrapper(tracer, is_text_completion) + ) + wrap_function_wrapper(module_name, func_name, wrapper) + + def _uninstrument(self, **kwargs): + for module_name, func_name, _, _ in _WRAPPED_METHODS: + try: + unwrap(module_name, func_name) + except Exception: + logger.debug("Failed to unwrap %s.%s", module_name, func_name) + + +def _build_sync_wrapper(tracer, is_text_completion): + def wrapper(wrapped, instance, args, kwargs): + return _invoke_completion( + tracer, + wrapped, + args, + kwargs, + is_text_completion=is_text_completion, + ) + + return wrapper + + +def _build_async_wrapper(tracer, is_text_completion): + async def wrapper(wrapped, instance, args, kwargs): + return await _invoke_acompletion( + tracer, + wrapped, + args, + kwargs, + is_text_completion=is_text_completion, + ) + + return wrapper + + +def _invoke_completion(tracer, wrapped, args, kwargs, *, is_text_completion=False): + if _should_skip_instrumentation(kwargs): + return wrapped(*args, **kwargs) + + span_name = _span_name(kwargs, is_text_completion) + request_type = _request_type(kwargs, is_text_completion) + attributes = { + GenAIAttributes.GEN_AI_SYSTEM: "litellm", + SpanAttributes.LLM_REQUEST_TYPE: request_type, + SpanAttributes.LLM_IS_STREAMING: False, + } + + span = tracer.start_span( + span_name, + kind=SpanKind.CLIENT, + attributes=attributes, + ) + with use_span(span, end_on_exit=False): + _set_request_attributes(span, args, kwargs, is_text_completion) + + updated_args, updated_kwargs = apply_prompt_safety( + span, args, kwargs, request_type, span_name + ) + _set_prompt_attributes( + span, + updated_args, + updated_kwargs, + request_type, + is_text_completion, + ) + + token = context_api.attach( + context_api.set_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY, True) + ) + try: + response = wrapped(*updated_args, **updated_kwargs) + except Exception as exc: + context_api.detach(token) + _record_span_error(span, exc) + span.end() + raise + + if inspect.isawaitable(response): + context_api.detach(token) + return _finalize_awaitable_response( + span, + response, + request_type, + span_name, + ) + + context_api.detach(token) + return _finalize_response(span, response, request_type, span_name) + + +async def _invoke_acompletion( + tracer, + wrapped, + args, + kwargs, + *, + is_text_completion=False, +): + if _should_skip_instrumentation(kwargs): + return await wrapped(*args, **kwargs) + + span_name = _span_name(kwargs, is_text_completion) + request_type = _request_type(kwargs, is_text_completion) + attributes = { + GenAIAttributes.GEN_AI_SYSTEM: "litellm", + SpanAttributes.LLM_REQUEST_TYPE: request_type, + SpanAttributes.LLM_IS_STREAMING: False, + } + + span = tracer.start_span( + span_name, + kind=SpanKind.CLIENT, + attributes=attributes, + ) + with use_span(span, end_on_exit=False): + _set_request_attributes(span, args, kwargs, is_text_completion) + + updated_args, updated_kwargs = apply_prompt_safety( + span, args, kwargs, request_type, span_name + ) + _set_prompt_attributes( + span, + updated_args, + updated_kwargs, + request_type, + is_text_completion, + ) + + token = context_api.attach( + context_api.set_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY, True) + ) + try: + response = await wrapped(*updated_args, **updated_kwargs) + except Exception as exc: + _record_span_error(span, exc) + span.end() + raise + finally: + context_api.detach(token) + + return _finalize_response(span, response, request_type, span_name) + + +def _should_skip_instrumentation(kwargs): + return bool( + context_api.get_value(_SUPPRESS_INSTRUMENTATION_KEY) + or context_api.get_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY) + or kwargs.get("stream") + ) + + +def _span_name(kwargs, is_text_completion): + if is_text_completion or kwargs.get("text_completion"): + return "litellm.text_completion" + return "litellm.completion" + + +def _request_type(kwargs, is_text_completion): + if is_text_completion or kwargs.get("text_completion"): + return LLMRequestTypeValues.COMPLETION.value + return LLMRequestTypeValues.CHAT.value + + +def _set_request_attributes(span, args, kwargs, is_text_completion): + model = kwargs.get("model") + if model is None and args: + model = args[1] if is_text_completion and len(args) > 1 else args[0] + if model is not None: + span.set_attribute(GenAIAttributes.GEN_AI_REQUEST_MODEL, str(model)) + + user = kwargs.get("user") + if user is not None: + span.set_attribute(SpanAttributes.LLM_USER, str(user)) + + custom_provider = kwargs.get("custom_llm_provider") + if custom_provider is not None: + span.set_attribute("litellm.request.provider", str(custom_provider)) + + +def _set_prompt_attributes(span, args, kwargs, request_type, is_text_completion): + if request_type == LLMRequestTypeValues.COMPLETION.value: + prompt = kwargs.get("prompt") + if prompt is None and args: + prompt = args[0] + + prompt_texts = extract_prompt_texts(prompt) + for index, text in enumerate(prompt_texts): + span.set_attribute(f"{SpanAttributes.LLM_PROMPTS}.{index}.role", "user") + span.set_attribute(f"{SpanAttributes.LLM_PROMPTS}.{index}.content", text) + return + + messages = kwargs.get("messages") + if messages is None and len(args) > 1: + messages = args[1] + if not isinstance(messages, list): + return + + for index, message in enumerate(messages): + role = _object_value(message, "role") + content = extract_text_content(_object_value(message, "content")) + if role is not None: + span.set_attribute(f"{SpanAttributes.LLM_PROMPTS}.{index}.role", str(role)) + if content: + span.set_attribute( + f"{SpanAttributes.LLM_PROMPTS}.{index}.content", + content, + ) + + +def _set_response_attributes(span, response): + response_model = _object_value(response, "model") + if response_model is not None: + span.set_attribute(GenAIAttributes.GEN_AI_RESPONSE_MODEL, str(response_model)) + + usage = _object_value(response, "usage") + input_tokens = _object_value(usage, "prompt_tokens") + output_tokens = _object_value(usage, "completion_tokens") + total_tokens = _object_value(usage, "total_tokens") + + if input_tokens is not None: + span.set_attribute(GenAIAttributes.GEN_AI_USAGE_INPUT_TOKENS, int(input_tokens)) + if output_tokens is not None: + span.set_attribute( + GenAIAttributes.GEN_AI_USAGE_OUTPUT_TOKENS, + int(output_tokens), + ) + if total_tokens is not None: + span.set_attribute(SpanAttributes.LLM_USAGE_TOTAL_TOKENS, int(total_tokens)) + + choices = _object_value(response, "choices") or [] + for index, choice in enumerate(choices): + finish_reason = _object_value(choice, "finish_reason") + if finish_reason is not None: + span.set_attribute( + f"{SpanAttributes.LLM_COMPLETIONS}.{index}.finish_reason", + str(finish_reason), + ) + + message = _object_value(choice, "message") + role = _object_value(message, "role") + content = extract_text_content(_object_value(message, "content")) + if role is not None: + span.set_attribute( + f"{SpanAttributes.LLM_COMPLETIONS}.{index}.role", + str(role), + ) + if content: + span.set_attribute( + f"{SpanAttributes.LLM_COMPLETIONS}.{index}.content", + content, + ) + continue + + text = _object_value(choice, "text") + if isinstance(text, str) and text: + span.set_attribute( + f"{SpanAttributes.LLM_COMPLETIONS}.{index}.content", + text, + ) + + +async def _finalize_awaitable_response( + span, + response, + request_type, + span_name, +): + token = context_api.attach( + context_api.set_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY, True) + ) + try: + awaited_response = await response + except Exception as exc: + _record_span_error(span, exc) + span.end() + raise + finally: + context_api.detach(token) + + return _finalize_response(span, awaited_response, request_type, span_name) + + +def _finalize_response(span, response, request_type, span_name): + try: + apply_completion_safety(span, response, request_type, span_name) + _set_response_attributes(span, response) + span.set_status(Status(StatusCode.OK)) + return response + finally: + span.end() + + +def _record_span_error(span, exc): + span.record_exception(exc) + span.set_status(Status(StatusCode.ERROR, str(exc))) + + +def _object_value(obj, key): + if obj is None: + return None + if isinstance(obj, dict): + return obj.get(key) + return getattr(obj, key, None) + + +__all__ = [ + "LiteLLMInstrumentor", + "_invoke_acompletion", + "_invoke_completion", +] diff --git a/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/safety.py b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/safety.py new file mode 100644 index 0000000000..7f12cdaaa1 --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/safety.py @@ -0,0 +1,445 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "LiteLLM" + + +def apply_prompt_safety(span, args, kwargs, request_type, span_name): + messages, source = _get_messages(args, kwargs) + if isinstance(messages, list): + return _apply_messages_prompt_safety( + span, + args, + kwargs, + messages, + source, + request_type, + span_name, + ) + + if request_type != LLMRequestTypeValues.COMPLETION.value: + return args, kwargs + + return _apply_text_prompt_safety(span, args, kwargs, request_type, span_name) + + +def _apply_messages_prompt_safety( + span, + args, + kwargs, + messages, + source, + request_type, + span_name, +): + updated_messages = messages + changed = False + + for index, message in enumerate(messages): + content = get_object_value(message, "content") + updated_content, content_changed = _mask_prompt_content( + span, + content, + span_name=span_name, + request_type=request_type, + segment_index=index, + segment_role=get_object_value(message, "role") or "user", + ) + if not content_changed: + continue + if updated_messages is messages: + updated_messages = clone_value(messages) + set_object_value(updated_messages[index], "content", updated_content) + changed = True + + if not changed: + return args, kwargs + + if source == "args": + updated_args = list(args) + updated_args[1] = updated_messages + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["messages"] = updated_messages + return args, updated_kwargs + + +def _apply_text_prompt_safety(span, args, kwargs, request_type, span_name): + prompt, source = _get_prompt(args, kwargs) + updated_prompt, changed = _mask_text_prompt_value( + span, + prompt, + span_name=span_name, + request_type=request_type, + ) + if not changed: + return args, kwargs + + if source == "args": + updated_args = list(args) + updated_args[0] = updated_prompt + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["prompt"] = updated_prompt + return args, updated_kwargs + + +def _mask_text_prompt_value(span, value, *, span_name, request_type): + if isinstance(value, str): + return _mask_prompt_text( + span, + value, + span_name=span_name, + request_type=request_type, + segment_index=0, + segment_role="user", + ) + + if not isinstance(value, list): + return value, False + + updated_value = value + for index, item in enumerate(value): + updated_item, changed = _mask_text_prompt_item( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[index] = updated_item + + return updated_value, updated_value is not value + + +def _mask_text_prompt_item( + span, + value, + *, + span_name, + request_type, + segment_index, + metadata=None, +): + if isinstance(value, str): + return _mask_prompt_text( + span, + value, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role="user", + metadata=metadata, + ) + + if not isinstance(value, list): + return value, False + + updated_value = value + for index, item in enumerate(value): + updated_item, changed = _mask_text_prompt_item( + span, + item, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + metadata={"nested_index": index, **(metadata or {})}, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[index] = updated_item + + return updated_value, updated_value is not value + + +def extract_prompt_texts(prompt): + texts = [] + _collect_prompt_texts(prompt, texts) + return texts + + +def _collect_prompt_texts(value, texts): + if isinstance(value, str): + texts.append(value) + return + + if not isinstance(value, list): + return + + for item in value: + _collect_prompt_texts(item, texts) + + +def _get_prompt(args, kwargs): + if "prompt" in kwargs: + return kwargs.get("prompt"), "kwargs" + if args: + return args[0], "args" + return None, None + + +def apply_completion_safety(span, response, request_type, span_name): + try: + choices = get_object_value(response, "choices") or [] + for index, choice in enumerate(choices): + message = get_object_value(choice, "message") + message_content = get_object_value(message, "content") if message is not None else None + if message is not None: + updated_content, content_changed = _mask_completion_content( + span, + message_content, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if content_changed: + set_object_value(message, "content", updated_content) + if isinstance(updated_content, str): + set_object_value(choice, "text", updated_content) + continue + + text = get_object_value(choice, "text") + if not isinstance(text, str): + continue + updated_text, text_changed = _mask_completion_text( + span, + text, + span_name=span_name, + request_type=request_type, + segment_index=index, + ) + if text_changed: + set_object_value(choice, "text", updated_text) + if message is not None and isinstance(message_content, str): + set_object_value(message, "content", updated_text) + except Exception: + return + + +def _get_messages(args, kwargs): + if len(args) > 1: + return args[1], "args" + return kwargs.get("messages"), "kwargs" + + +def _mask_prompt_content( + span, + content, + *, + span_name, + request_type, + segment_index, + segment_role, + metadata=None, +): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_prompt_text( + span, + block, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text", "input_text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content( + span, + content, + *, + span_name, + request_type, + segment_index, + metadata=None, +): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + metadata=metadata, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_completion_text( + span, + block, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text", "output_text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + metadata={"block_index": block_index, **(metadata or {})}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + span_name, + request_type, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text( + span, + text, + *, + span_name, + request_type, + segment_index, + metadata=None, +): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role="assistant", + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def extract_text_content(content): + if isinstance(content, str): + return content + + if not isinstance(content, list): + return None + + parts = [] + for block in content: + if isinstance(block, str): + parts.append(block) + continue + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type in (None, "text", "input_text", "output_text") and isinstance( + block_text, str + ): + parts.append(block_text) + + if not parts: + return None + return "\n".join(parts) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/version.py b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/version.py new file mode 100644 index 0000000000..c70d52c4be --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm/version.py @@ -0,0 +1 @@ +__version__ = "0.52.6" diff --git a/packages/opentelemetry-instrumentation-litellm/project.json b/packages/opentelemetry-instrumentation-litellm/project.json new file mode 100644 index 0000000000..621a854106 --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/project.json @@ -0,0 +1,77 @@ +{ + "name": "opentelemetry-instrumentation-litellm", + "$schema": "../../node_modules/nx/schemas/project-schema.json", + "projectType": "library", + "sourceRoot": "packages/opentelemetry-instrumentation-litellm/opentelemetry/instrumentation/litellm", + "targets": { + "lock": { + "executor": "nx:run-commands", + "options": { + "command": "uv lock", + "cwd": "packages/opentelemetry-instrumentation-litellm" + } + }, + "add": { + "executor": "@nxlv/python:add", + "options": {} + }, + "update": { + "executor": "@nxlv/python:update", + "options": {} + }, + "remove": { + "executor": "@nxlv/python:remove", + "options": {} + }, + "build": { + "executor": "@nxlv/python:build", + "outputs": [ + "{projectRoot}/dist" + ], + "options": { + "outputPath": "packages/opentelemetry-instrumentation-litellm/dist", + "publish": false, + "lockedVersions": true, + "bundleLocalDependencies": true + } + }, + "install": { + "executor": "nx:run-commands", + "options": { + "command": "uv sync --all-groups", + "cwd": "packages/opentelemetry-instrumentation-litellm" + } + }, + "lint": { + "executor": "nx:run-commands", + "options": { + "command": "uv run ruff check .", + "cwd": "packages/opentelemetry-instrumentation-litellm" + } + }, + "test": { + "executor": "nx:run-commands", + "outputs": [ + "{workspaceRoot}/reports/packages/opentelemetry-instrumentation-litellm/unittests", + "{workspaceRoot}/coverage/packages/opentelemetry-instrumentation-litellm" + ], + "options": { + "command": "uv run pytest tests/", + "cwd": "packages/opentelemetry-instrumentation-litellm" + } + }, + "build-release": { + "executor": "nx:run-commands", + "options": { + "commands": [ + "chmod +x ../../scripts/build-release.sh", + "../../scripts/build-release.sh" + ], + "cwd": "packages/opentelemetry-instrumentation-litellm" + } + } + }, + "tags": [ + "instrumentation" + ] +} diff --git a/packages/opentelemetry-instrumentation-litellm/pyproject.toml b/packages/opentelemetry-instrumentation-litellm/pyproject.toml new file mode 100644 index 0000000000..30aedfda2a --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/pyproject.toml @@ -0,0 +1,81 @@ +[project] +name = "opentelemetry-instrumentation-litellm" +version = "0.52.6" +description = "OpenTelemetry LiteLLM instrumentation" +authors = [ + { name = "FortifyRoot", email = "engineering@fortifyroot.com" }, +] +license = "Apache-2.0" +readme = "README.md" +requires-python = ">=3.10,<4" +dependencies = [ + "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", + "opentelemetry-instrumentation>=0.59b0", + "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", + "opentelemetry-semantic-conventions>=0.59b0", +] + +[project.urls] +Repository = "https://github.com/traceloop/openllmetry/tree/main/packages/opentelemetry-instrumentation-litellm" + +[project.optional-dependencies] +instruments = ["litellm>=1.71.2,<2"] + +[project.entry-points."opentelemetry_instrumentor"] +litellm = "opentelemetry.instrumentation.litellm:LiteLLMInstrumentor" + +[dependency-groups] +dev = [ + "autopep8>=2.2.0,<3", + "black>=25.1.0,<26", + "isort>=6.0.1,<7", + "ruff>=0.4.0", +] +test = [ + "litellm>=1.71.2,<2", + "opentelemetry-sdk>=1.38.0,<2", + "pytest-asyncio>=0.23.7,<0.24.0", + "pytest-sugar==1.0.0", + "pytest>=8.2.2,<9", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build.targets.wheel] +packages = ["opentelemetry"] + +[tool.coverage.run] +branch = true +source = ["opentelemetry/instrumentation/litellm"] + +[tool.coverage.report] +exclude_lines = ["if TYPE_CHECKING:"] +show_missing = true + +[tool.ruff] +line-length = 120 +exclude = [ + ".git", + "__pycache__", + "build", + "dist", + ".venv", + ".pytest_cache", +] + +[tool.ruff.lint] +select = ["E", "F", "W"] + +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + +[tool.uv] +constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-litellm/tests/__init__.py b/packages/opentelemetry-instrumentation-litellm/tests/__init__.py new file mode 100644 index 0000000000..8b13789179 --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/tests/__init__.py @@ -0,0 +1 @@ + diff --git a/packages/opentelemetry-instrumentation-litellm/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-litellm/tests/test_safety_hooks.py new file mode 100644 index 0000000000..8b7f6cacd7 --- /dev/null +++ b/packages/opentelemetry-instrumentation-litellm/tests/test_safety_hooks.py @@ -0,0 +1,481 @@ +from __future__ import annotations + +import inspect +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +from litellm.types.utils import ModelResponse, TextCompletionResponse +from opentelemetry import context as context_api +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.litellm import ( + LiteLLMInstrumentor, + _WRAPPED_METHODS, + _invoke_acompletion, + _invoke_completion, +) +from opentelemetry.instrumentation.litellm.safety import ( + apply_completion_safety, + apply_prompt_safety, + extract_prompt_texts, + extract_text_content, + _get_messages, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv._incubating.attributes import ( + gen_ai_attributes as GenAIAttributes, +) +from opentelemetry.semconv_ai import ( + SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY, + SpanAttributes, +) + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_tracer(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def _prompt_result(masked_text, context): + return SafetyResult( + text=masked_text, + overall_action="MASK", + findings=[ + SafetyFinding( + category="PII", + severity="HIGH", + action="MASK", + rule_name="PII.secret", + start=0, + end=len(context.text), + ) + ], + ) + + +def _completion_result(masked_text, context): + return SafetyResult( + text=masked_text, + overall_action="MASK", + findings=[ + SafetyFinding( + category="SECRET", + severity="HIGH", + action="MASK", + rule_name="SECRET.token", + start=0, + end=len(context.text), + ) + ], + ) + + +def test_sync_completion_masks_prompt_response_and_sets_span_attributes(): + exporter, tracer = _test_tracer() + register_prompt_safety_handler( + lambda context: _prompt_result("[PII.email]", context) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: _completion_result("[SECRET.token]", context) + if context.location == SafetyLocation.COMPLETION and context.text == "token-123" + else None + ) + + messages = [{"role": "user", "content": "secret"}] + + def wrapped(*args, **kwargs): + assert kwargs["messages"][0]["content"] == "[PII.email]" + assert context_api.get_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY) is True + return ModelResponse( + model="gpt-4o-mini", + usage={"prompt_tokens": 3, "completion_tokens": 5, "total_tokens": 8}, + choices=[ + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "token-123"}, + } + ], + ) + + response = _invoke_completion( + tracer, + wrapped, + (), + {"model": "gpt-4o-mini", "messages": messages}, + ) + + assert messages[0]["content"] == "secret" + assert response.choices[0].message.content == "[SECRET.token]" + + spans = exporter.get_finished_spans() + assert len(spans) == 1 + span = spans[0] + assert span.name == "litellm.completion" + assert span.attributes[GenAIAttributes.GEN_AI_REQUEST_MODEL] == "gpt-4o-mini" + assert span.attributes[GenAIAttributes.GEN_AI_RESPONSE_MODEL] == "gpt-4o-mini" + assert span.attributes[f"{SpanAttributes.LLM_PROMPTS}.0.content"] == "[PII.email]" + assert span.attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[SECRET.token]" + assert span.attributes[SpanAttributes.LLM_USAGE_TOTAL_TOKENS] == 8 + assert len(span.events) == 2 + + +def test_sync_text_completion_masks_text_choices(): + exporter, tracer = _test_tracer() + register_prompt_safety_handler( + lambda context: _prompt_result("[PII.email]", context) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: _completion_result("[SECRET.token]", context) + if context.location == SafetyLocation.COMPLETION and context.text == "token-123" + else None + ) + + def wrapped(*args, **kwargs): + assert args[0] == "[PII.email]" + assert kwargs["model"] == "gpt-3.5-turbo-instruct" + assert context_api.get_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY) is True + return TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + choices=[{"text": "token-123"}], + ) + + response = _invoke_completion( + tracer, + wrapped, + ("secret",), + {"model": "gpt-3.5-turbo-instruct"}, + is_text_completion=True, + ) + + assert response.choices[0].text == "[SECRET.token]" + + spans = exporter.get_finished_spans() + assert len(spans) == 1 + span = spans[0] + assert span.name == "litellm.text_completion" + assert span.attributes[SpanAttributes.LLM_REQUEST_TYPE] == "completion" + assert span.attributes[GenAIAttributes.GEN_AI_REQUEST_MODEL] == "gpt-3.5-turbo-instruct" + assert span.attributes[f"{SpanAttributes.LLM_PROMPTS}.0.content"] == "[PII.email]" + assert span.attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[SECRET.token]" + + +@pytest.mark.asyncio +async def test_async_completion_masks_prompt_and_response(): + exporter, tracer = _test_tracer() + register_prompt_safety_handler( + lambda context: _prompt_result("[PII.email]", context) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: _completion_result("[SECRET.token]", context) + if context.location == SafetyLocation.COMPLETION and context.text == "token-123" + else None + ) + + async def wrapped(*args, **kwargs): + assert kwargs["messages"][0]["content"] == "[PII.email]" + return ModelResponse( + model="gpt-4o-mini", + choices=[ + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "token-123"}, + } + ], + ) + + response = await _invoke_acompletion( + tracer, + wrapped, + (), + {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "secret"}]}, + ) + + assert response.choices[0].message.content == "[SECRET.token]" + spans = exporter.get_finished_spans() + assert len(spans) == 1 + assert spans[0].attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[SECRET.token]" + + +@pytest.mark.asyncio +async def test_async_text_completion_masks_prompt_and_response(): + exporter, tracer = _test_tracer() + register_prompt_safety_handler( + lambda context: _prompt_result("[PII.email]", context) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: _completion_result("[SECRET.token]", context) + if context.location == SafetyLocation.COMPLETION and context.text == "token-123" + else None + ) + + async def wrapped(*args, **kwargs): + assert args[0] == "[PII.email]" + assert kwargs["model"] == "gpt-3.5-turbo-instruct" + return TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + choices=[{"text": "token-123"}], + ) + + response = await _invoke_acompletion( + tracer, + wrapped, + ("secret",), + {"model": "gpt-3.5-turbo-instruct"}, + is_text_completion=True, + ) + + assert response.choices[0].text == "[SECRET.token]" + span = exporter.get_finished_spans()[0] + assert span.name == "litellm.text_completion" + assert span.attributes[f"{SpanAttributes.LLM_PROMPTS}.0.content"] == "[PII.email]" + assert span.attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[SECRET.token]" + + +@pytest.mark.asyncio +async def test_sync_wrapper_handles_awaitable_response(): + exporter, tracer = _test_tracer() + register_prompt_safety_handler( + lambda context: _prompt_result("[PII.email]", context) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: _completion_result("[SECRET.token]", context) + if context.location == SafetyLocation.COMPLETION and context.text == "token-123" + else None + ) + + async def response_coro(): + assert context_api.get_value(SUPPRESS_LANGUAGE_MODEL_INSTRUMENTATION_KEY) is True + return ModelResponse( + model="gpt-4o-mini", + choices=[ + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "token-123"}, + } + ], + ) + + def wrapped(*args, **kwargs): + assert kwargs["messages"][0]["content"] == "[PII.email]" + return response_coro() + + response = _invoke_completion( + tracer, + wrapped, + (), + {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "secret"}]}, + ) + + assert inspect.iscoroutine(response) + response = await response + + assert response.choices[0].message.content == "[SECRET.token]" + spans = exporter.get_finished_spans() + assert len(spans) == 1 + assert spans[0].attributes[f"{SpanAttributes.LLM_COMPLETIONS}.0.content"] == "[SECRET.token]" + + +def test_streaming_completion_is_passthrough(): + exporter, tracer = _test_tracer() + sentinel = SimpleNamespace(value="stream") + calls = [] + + def wrapped(*args, **kwargs): + calls.append(kwargs) + return sentinel + + response = _invoke_completion( + tracer, + wrapped, + (), + { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "secret"}], + "stream": True, + }, + ) + + assert response is sentinel + assert calls[0]["messages"][0]["content"] == "secret" + assert not exporter.get_finished_spans() + + +def test_instrumentor_wraps_and_unwraps_all_methods(): + instrumentor = LiteLLMInstrumentor() + unwrap_side_effect = [None, RuntimeError("boom")] + [None] * (len(_WRAPPED_METHODS) - 2) + + with patch("opentelemetry.instrumentation.litellm.wrap_function_wrapper") as wrap_mock, patch( + "opentelemetry.instrumentation.litellm.unwrap", + side_effect=unwrap_side_effect, + ) as unwrap_mock, patch("opentelemetry.instrumentation.litellm.logger.debug") as debug_mock: + instrumentor._instrument() + instrumentor._uninstrument() + + assert instrumentor.instrumentation_dependencies() == ("litellm >= 1.71.2, < 2",) + assert wrap_mock.call_count == len(_WRAPPED_METHODS) + assert unwrap_mock.call_count == len(_WRAPPED_METHODS) + debug_mock.assert_called_once() + + +def test_error_paths_record_error_status_and_request_attributes(): + exporter, tracer = _test_tracer() + + def wrapped(*args, **kwargs): + raise ValueError("sync boom") + + with pytest.raises(ValueError, match="sync boom"): + _invoke_completion( + tracer, + wrapped, + ("gpt-4o-mini",), + {"user": "alice", "custom_llm_provider": "openai"}, + ) + + span = exporter.get_finished_spans()[0] + assert span.status.status_code.name == "ERROR" + assert span.attributes[GenAIAttributes.GEN_AI_REQUEST_MODEL] == "gpt-4o-mini" + assert span.attributes[SpanAttributes.LLM_USER] == "alice" + assert span.attributes["litellm.request.provider"] == "openai" + + +@pytest.mark.asyncio +async def test_async_skip_and_error_paths(): + exporter, tracer = _test_tracer() + sentinel = SimpleNamespace(value="stream") + + async def skipped(*args, **kwargs): + return sentinel + + response = await _invoke_acompletion( + tracer, + skipped, + (), + {"model": "gpt-4o-mini", "stream": True}, + ) + assert response is sentinel + assert not exporter.get_finished_spans() + + async def raising(*args, **kwargs): + raise RuntimeError("async boom") + + with pytest.raises(RuntimeError, match="async boom"): + await _invoke_acompletion( + tracer, + raising, + (), + {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "secret"}]}, + ) + + error_span = exporter.get_finished_spans()[0] + assert error_span.status.status_code.name == "ERROR" + + +def test_safety_helpers_cover_args_blocks_and_text_extraction(): + register_prompt_safety_handler( + lambda context: _prompt_result(f"[MASKED:{context.text}]", context) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: _completion_result(f"[MASKED:{context.text}]", context) + if context.location == SafetyLocation.COMPLETION + else None + ) + + args = ( + "gpt-4o-mini", + [ + { + "role": "user", + "content": [ + "prompt-a", + {"type": "input_text", "text": "prompt-b"}, + {"type": "image_url", "text": "ignored"}, + ], + } + ], + ) + updated_args, unchanged_kwargs = apply_prompt_safety( + None, args, {}, "chat", "litellm.completion" + ) + assert unchanged_kwargs == {} + assert updated_args[1][0]["content"][:2] == ["[MASKED:prompt-a]", {"type": "input_text", "text": "[MASKED:prompt-b]"}] + assert _get_messages(updated_args, {})[1] == "args" + assert apply_prompt_safety(None, (), {"messages": "invalid"}, "chat", "litellm.completion") == ( + (), + {"messages": "invalid"}, + ) + unchanged_prompt_args, updated_prompt_kwargs = apply_prompt_safety( + None, + (), + {"prompt": ["prompt-a", ["prompt-b", [1, 2, 3]]]}, + "completion", + "litellm.text_completion", + ) + assert unchanged_prompt_args == () + assert updated_prompt_kwargs["prompt"][0] == "[MASKED:prompt-a]" + assert updated_prompt_kwargs["prompt"][1][0] == "[MASKED:prompt-b]" + + response = SimpleNamespace( + choices=[ + SimpleNamespace( + message=SimpleNamespace( + content=["completion-a", {"type": "output_text", "text": "completion-b"}] + ) + ), + SimpleNamespace( + message=SimpleNamespace(content="completion-c"), + text="completion-c", + ), + SimpleNamespace(text="fallback-text"), + ] + ) + apply_completion_safety(None, response, "chat", "litellm.completion") + assert response.choices[0].message.content[:2] == [ + "[MASKED:completion-a]", + {"type": "output_text", "text": "[MASKED:completion-b]"}, + ] + assert response.choices[1].message.content == "[MASKED:completion-c]" + assert response.choices[1].text == "[MASKED:completion-c]" + assert response.choices[2].text == "[MASKED:fallback-text]" + assert extract_prompt_texts(["a", ["b", [1, 2, 3]]]) == ["a", "b"] + assert extract_text_content(["a", {"type": "output_text", "text": "b"}]) == "a\nb" + assert extract_text_content([{"type": "image_url", "text": "ignored"}]) is None + assert _resolve_masked_text("same", None) == ("same", False) + assert _resolve_masked_text("same", SafetyResult(text="same", overall_action="MASK", findings=[])) == ( + "same", + False, + ) diff --git a/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/dispatcher_wrapper.py b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/dispatcher_wrapper.py index 2e80cafd52..2fa0d5ebbe 100644 --- a/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/dispatcher_wrapper.py +++ b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/dispatcher_wrapper.py @@ -24,6 +24,7 @@ LLMChatEndEvent, LLMChatStartEvent, LLMCompletionEndEvent, + LLMCompletionStartEvent, LLMPredictEndEvent, ) from llama_index.core.instrumentation.events.rerank import ReRankStartEvent @@ -35,6 +36,13 @@ emit_chat_response_events, emit_rerank_message_event, ) +from opentelemetry.instrumentation.llamaindex.safety import ( + apply_chat_end_safety, + apply_completion_end_safety, + apply_completion_start_attributes, + apply_predict_end_safety, + instrument_llm_safety_wrappers, +) from opentelemetry.instrumentation.llamaindex.span_utils import ( set_embedding, set_llm_chat_request, @@ -71,6 +79,7 @@ def instrument_with_dispatcher(tracer: Tracer): + instrument_llm_safety_wrappers() dispatcher = get_dispatcher() openllmetry_span_handler = OpenLLMetrySpanHandler(tracer) dispatcher.add_span_handler(openllmetry_span_handler) @@ -127,14 +136,24 @@ def _(self, event: LLMChatStartEvent): @update_span_for_event.register def _(self, event: LLMChatEndEvent): + apply_chat_end_safety(event, self.otel_span) set_llm_chat_response_model_attributes(event, self.otel_span) if should_emit_events(): emit_chat_response_events(event) else: set_llm_chat_response(event, self.otel_span) # noqa: F821 + @update_span_for_event.register + def _(self, event: LLMCompletionStartEvent): + apply_completion_start_attributes(event, self.otel_span) + + @update_span_for_event.register + def _(self, event: LLMCompletionEndEvent): + apply_completion_end_safety(event, self.otel_span) + @update_span_for_event.register def _(self, event: LLMPredictEndEvent): + apply_predict_end_safety(event, self.otel_span) if not should_emit_events(): set_llm_predict_response(event, self.otel_span) diff --git a/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py new file mode 100644 index 0000000000..023128a296 --- /dev/null +++ b/packages/opentelemetry-instrumentation-llamaindex/opentelemetry/instrumentation/llamaindex/safety.py @@ -0,0 +1,449 @@ +from __future__ import annotations + +import importlib +import pkgutil +from typing import Any + +import llama_index.core.llms +from llama_index.core.base.llms.base import BaseLLM +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.instrumentation.llamaindex.utils import should_send_prompts +from opentelemetry.semconv._incubating.attributes import ( + gen_ai_attributes as GenAIAttributes, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues, SpanAttributes +from wrapt import wrap_function_wrapper + +try: + import llama_index.llms +except Exception: # pragma: no cover + llama_index = None + +PROVIDER = "LlamaIndex" +_WRAPPERS_INSTALLED = False +_METHODS = ("chat", "achat", "complete", "acomplete") + + +def instrument_llm_safety_wrappers(): + global _WRAPPERS_INSTALLED + if _WRAPPERS_INSTALLED: + return + + _wrap_base_methods() + + for package_name in ("llama_index.core.llms", "llama_index.llms"): + _wrap_package_classes(package_name) + + _WRAPPERS_INSTALLED = True + + +def llm_chat_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_chat_prompt_safety(instance, args, kwargs) + return wrapped(*updated_args, **updated_kwargs) + + +async def llm_achat_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_chat_prompt_safety(instance, args, kwargs) + return await wrapped(*updated_args, **updated_kwargs) + + +def llm_complete_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_completion_prompt_safety(instance, args, kwargs) + return wrapped(*updated_args, **updated_kwargs) + + +async def llm_acomplete_wrapper(wrapped, instance, args, kwargs): + updated_args, updated_kwargs = _apply_completion_prompt_safety(instance, args, kwargs) + return await wrapped(*updated_args, **updated_kwargs) + + +_METHOD_WRAPPERS = { + "chat": llm_chat_wrapper, + "achat": llm_achat_wrapper, + "complete": llm_complete_wrapper, + "acomplete": llm_acomplete_wrapper, +} + + +def apply_chat_end_safety(event, span): + _apply_chat_response_safety(event, span) + + +def apply_completion_start_attributes(event, span): + if span is None or not span.is_recording(): + return + + model_dict = event.model_dict or {} + if "llm" in model_dict: + model_dict = model_dict.get("llm", {}) + + span.set_attribute( + SpanAttributes.LLM_REQUEST_TYPE, LLMRequestTypeValues.COMPLETION.value + ) + span.set_attribute(GenAIAttributes.GEN_AI_REQUEST_MODEL, model_dict.get("model")) + span.set_attribute( + GenAIAttributes.GEN_AI_REQUEST_TEMPERATURE, + model_dict.get("temperature"), + ) + + if should_send_prompts(): + span.set_attribute(f"{GenAIAttributes.GEN_AI_PROMPT}.0.role", "user") + span.set_attribute(f"{GenAIAttributes.GEN_AI_PROMPT}.0.content", event.prompt) + + +def apply_completion_end_safety(event, span): + if span is None or not span.is_recording(): + return + + response = event.response + if response is not None and isinstance(get_object_value(response, "text"), str): + updated_text, changed = _mask_completion_text( + get_object_value(response, "text"), + span=span, + span_name=getattr(span, "name", "llamaindex.completion"), + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=0, + segment_role="assistant", + ) + if changed: + set_object_value(response, "text", updated_text) + + _set_completion_response_model_attributes(response, span) + + if should_send_prompts(): + span.set_attribute(f"{GenAIAttributes.GEN_AI_COMPLETION}.0.role", "assistant") + span.set_attribute( + f"{GenAIAttributes.GEN_AI_COMPLETION}.0.content", + get_object_value(response, "text"), + ) + + +def apply_predict_end_safety(event, span): + if not isinstance(event.output, str): + return + + updated_text, changed = _mask_completion_text( + event.output, + span=span, + span_name=getattr(span, "name", "llamaindex.predict"), + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=0, + segment_role="assistant", + ) + if changed: + event.output = updated_text + + +def _wrap_package_classes(package_name: str): + try: + package = importlib.import_module(package_name) + except Exception: + return + + for module_info in pkgutil.iter_modules(package.__path__): + module_name = f"{package_name}.{module_info.name}" + try: + module = importlib.import_module(module_name) + except Exception: + continue + for _, cls in module.__dict__.items(): + if not isinstance(cls, type): + continue + if not issubclass(cls, BaseLLM): + continue + for method_name in _METHODS: + if method_name not in cls.__dict__: + continue + wrap_function_wrapper( + cls.__module__, + f"{cls.__name__}.{method_name}", + _METHOD_WRAPPERS[method_name], + ) + + +def _wrap_base_methods(): + for method_name in _METHODS: + wrap_function_wrapper( + "llama_index.core.base.llms.base", + f"BaseLLM.{method_name}", + _METHOD_WRAPPERS[method_name], + ) + + +def _apply_chat_prompt_safety(instance, args, kwargs): + messages = args[0] if args else kwargs.get("messages") + if not isinstance(messages, (list, tuple)): + return args, kwargs + + updated_messages = list(clone_value(list(messages))) + changed = False + span_name = f"{instance.__class__.__name__}.chat" + + for index, message in enumerate(updated_messages): + if _mask_chat_message( + message, + span=None, + span_name=span_name, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=index, + ): + changed = True + + if not changed: + return args, kwargs + + if args: + updated_args = list(args) + updated_args[0] = updated_messages + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["messages"] = updated_messages + return args, updated_kwargs + + +def _apply_completion_prompt_safety(instance, args, kwargs): + prompt = args[0] if args else kwargs.get("prompt") + if not isinstance(prompt, str): + return args, kwargs + + updated_prompt, changed = _mask_prompt_text( + prompt, + span=None, + span_name=f"{instance.__class__.__name__}.completion", + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=0, + segment_role="user", + ) + if not changed: + return args, kwargs + + if args: + updated_args = list(args) + updated_args[0] = updated_prompt + return tuple(updated_args), kwargs + + updated_kwargs = dict(kwargs) + updated_kwargs["prompt"] = updated_prompt + return args, updated_kwargs + + +def _apply_chat_response_safety(event, span): + response = event.response + if response is None: + return + + message = get_object_value(response, "message") + if message is None: + return + + _mask_chat_message( + message, + span=span, + span_name=getattr(span, "name", "llamaindex.chat"), + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=0, + segment_role="assistant", + ) + + +def _mask_chat_message( + message, + *, + span, + span_name, + request_type, + segment_index, + segment_role=None, +): + changed = False + role = segment_role or _message_role(message) + blocks = get_object_value(message, "blocks") + if isinstance(blocks, list): + for block_index, block in enumerate(blocks): + text = _block_text(block) + if not isinstance(text, str): + continue + updated_text, text_changed = _mask_text( + text, + span=span, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=role, + location=SafetyLocation.PROMPT if span is None else SafetyLocation.COMPLETION, + metadata={"block_index": block_index}, + ) + if not text_changed: + continue + _set_block_text(block, updated_text) + changed = True + return changed + + content = get_object_value(message, "content") + if not isinstance(content, str): + return False + updated_text, text_changed = _mask_text( + content, + span=span, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=role, + location=SafetyLocation.PROMPT if span is None else SafetyLocation.COMPLETION, + ) + if text_changed: + set_object_value(message, "content", updated_text) + return True + return False + + +def _set_completion_response_model_attributes(response, span): + if response is None: + return + + raw = get_object_value(response, "raw") + if raw is None: + return + + model = get_object_value(raw, "model") + if model: + span.set_attribute(GenAIAttributes.GEN_AI_RESPONSE_MODEL, model) + + usage = get_object_value(raw, "usage") + if usage is not None: + completion_tokens = get_object_value(usage, "completion_tokens") + prompt_tokens = get_object_value(usage, "prompt_tokens") + total_tokens = get_object_value(usage, "total_tokens") + if completion_tokens is not None: + span.set_attribute( + GenAIAttributes.GEN_AI_USAGE_OUTPUT_TOKENS, + int(completion_tokens), + ) + if prompt_tokens is not None: + span.set_attribute( + GenAIAttributes.GEN_AI_USAGE_INPUT_TOKENS, + int(prompt_tokens), + ) + if total_tokens is not None: + span.set_attribute( + SpanAttributes.LLM_USAGE_TOTAL_TOKENS, + int(total_tokens), + ) + + +def _mask_prompt_text( + text, + *, + span, + span_name, + request_type, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text( + text, + *, + span, + span_name, + request_type, + segment_index, + segment_role, +): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_text( + text, + *, + span, + span_name, + request_type, + segment_index, + segment_role, + location, + metadata=None, +): + if location == SafetyLocation.PROMPT: + return _mask_prompt_text( + text, + span=span, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _mask_completion_text( + text, + span=span, + span_name=span_name, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + + +def _block_text(block): + block_type = get_object_value(block, "block_type") + if block_type == "text": + return get_object_value(block, "text") + if block_type == "thinking": + return get_object_value(block, "content") + return None + + +def _set_block_text(block, value): + block_type = get_object_value(block, "block_type") + if block_type == "text": + return set_object_value(block, "text", value) + if block_type == "thinking": + return set_object_value(block, "content", value) + return False + + +def _message_role(message) -> str: + role = get_object_value(message, "role") + role_value = getattr(role, "value", role) + return str(role_value or "user").lower() + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-llamaindex/pyproject.toml b/packages/opentelemetry-instrumentation-llamaindex/pyproject.toml index ac125dc306..57aa18dfa5 100644 --- a/packages/opentelemetry-instrumentation-llamaindex/pyproject.toml +++ b/packages/opentelemetry-instrumentation-llamaindex/pyproject.toml @@ -13,6 +13,7 @@ requires-python = ">=3.10,<4" dependencies = [ "inflection>=0.5.1,<0.6.0", "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -84,3 +85,9 @@ select = ["E", "F", "W"] [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } + +[tool.pytest.ini_options] +markers = ["safety: tests for FortifyRoot safety masking hooks"] diff --git a/packages/opentelemetry-instrumentation-llamaindex/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-llamaindex/tests/test_safety_hooks.py new file mode 100644 index 0000000000..47de958073 --- /dev/null +++ b/packages/opentelemetry-instrumentation-llamaindex/tests/test_safety_hooks.py @@ -0,0 +1,388 @@ +from llama_index.core.base.llms.types import ( + ChatMessage, + ChatResponse, + CompletionResponse, + MessageRole, +) +from llama_index.core.instrumentation.events.llm import ( + LLMChatEndEvent, + LLMCompletionEndEvent, + LLMCompletionStartEvent, +) +from types import SimpleNamespace +from unittest.mock import patch +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, +) + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.llamaindex.safety import ( + _apply_chat_prompt_safety, + _apply_completion_prompt_safety, + _block_text, + _set_block_text, + apply_chat_end_safety, + apply_completion_end_safety, + apply_completion_start_attributes, + apply_predict_end_safety, + instrument_llm_safety_wrappers, + llm_acomplete_wrapper, + llm_achat_wrapper, + llm_complete_wrapper, + llm_chat_wrapper, + _mask_chat_message, + _resolve_masked_text, +) + +pytestmark = pytest.mark.safety + + +class _FakeLLM: + pass + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_chat_prompt_safety_masks_messages_before_wrapped_call(): + captured = {} + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.llamaindex]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="PII", + severity="MEDIUM", + action="MASK", + rule_name="PII.llamaindex", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + def wrapped(messages, **kwargs): + captured["messages"] = messages + return "ok" + + messages = [ChatMessage(content="secret", role=MessageRole.USER)] + result = llm_chat_wrapper(wrapped, _FakeLLM(), (messages,), {}) + + assert result == "ok" + assert captured["messages"][0].content == "[PII.llamaindex]" + assert messages[0].content == "secret" + + +def test_completion_prompt_safety_masks_prompt_before_wrapped_call(): + captured = {} + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.llamaindex]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="PII", + severity="MEDIUM", + action="MASK", + rule_name="PII.llamaindex", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + def wrapped(prompt, **kwargs): + captured["prompt"] = prompt + return CompletionResponse(text="ok") + + llm_complete_wrapper(wrapped, _FakeLLM(), ("secret",), {}) + + assert captured["prompt"] == "[PII.llamaindex]" + + +def test_completion_start_sets_masked_prompt_attributes(): + exporter, tracer = _test_span() + + with tracer.start_as_current_span("llamaindex.completion") as span: + apply_completion_start_attributes( + LLMCompletionStartEvent( + prompt="[PII.llamaindex]", + additional_kwargs={}, + model_dict={"model": "test-model", "temperature": 0.1}, + span_id="span-1", + ), + span, + ) + + spans = exporter.get_finished_spans() + assert spans[0].attributes["gen_ai.prompt.0.content"] == "[PII.llamaindex]" + assert spans[0].attributes["gen_ai.prompt.0.role"] == "user" + + +def test_completion_end_safety_masks_response_and_span_attributes(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.llamaindex]", + overall_action="MASK", + findings=[ + SafetyFinding( + category="SECRET", + severity="HIGH", + action="MASK", + rule_name="SECRET.llamaindex", + start=0, + end=len(context.text), + ) + ], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = CompletionResponse(text="secret") + with tracer.start_as_current_span("llamaindex.completion") as span: + apply_completion_end_safety( + LLMCompletionEndEvent(prompt="prompt", response=response, span_id="span-1"), + span, + ) + + spans = exporter.get_finished_spans() + assert response.text == "[SECRET.llamaindex]" + assert spans[0].attributes["gen_ai.completion.0.content"] == "[SECRET.llamaindex]" + + +@pytest.mark.asyncio +async def test_async_wrappers_and_chat_end_safety_cover_block_paths(): + exporter, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "MEDIUM", "MASK", "PII.llamaindex", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.llamaindex", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + async def wrapped_messages(messages, **kwargs): + return messages + + async def wrapped_prompt(prompt, **kwargs): + return CompletionResponse(text="ok") + + messages = [ + SimpleNamespace( + role=MessageRole.USER, + blocks=[SimpleNamespace(block_type="text", text="secret")], + ) + ] + masked_messages = await llm_achat_wrapper(wrapped_messages, _FakeLLM(), (messages,), {}) + await llm_acomplete_wrapper(wrapped_prompt, _FakeLLM(), ("secret",), {}) + + with tracer.start_as_current_span("llamaindex.chat") as span: + _mask_chat_message( + messages[0], + span=span, + span_name="llamaindex.chat", + request_type="chat", + segment_index=0, + segment_role="assistant", + ) + response = ChatResponse(message=ChatMessage(role=MessageRole.ASSISTANT, content="secret")) + apply_chat_end_safety( + LLMChatEndEvent(messages=[], response=response, span_id="span-1"), + span, + ) + + assert masked_messages[0].blocks[0].text == "[MASKED:secret]" + assert messages[0].blocks[0].text == "[MASKED:secret]" + assert response.message.content == "[MASKED:secret]" + assert exporter.get_finished_spans() + + +def test_predict_and_registration_helpers_cover_non_emit_paths(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.llamaindex", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + event = SimpleNamespace(output="secret") + with patch( + "opentelemetry.instrumentation.llamaindex.safety.wrap_function_wrapper" + ) as wrap_mock, patch( + "opentelemetry.instrumentation.llamaindex.safety.pkgutil.iter_modules", + return_value=[SimpleNamespace(name="fake_module")], + ), patch( + "opentelemetry.instrumentation.llamaindex.safety.importlib.import_module" + ) as import_mock, patch( + "opentelemetry.instrumentation.llamaindex.safety._WRAPPERS_INSTALLED", False + ): + fake_module = SimpleNamespace() + + class FakeWrappedLLM(_FakeLLM.__class__): + pass + + class DummyLLM: + def chat(self): # pragma: no cover - signature only + return None + + def complete(self): # pragma: no cover - signature only + return None + + import llama_index.core.base.llms.base as base_module + + DummyLLM = type( + "DummyLLM", + (base_module.BaseLLM,), + { + "__module__": "fake.module", + "chat": lambda self: None, + "complete": lambda self: None, + }, + ) + fake_module.DummyLLM = DummyLLM + import_mock.side_effect = [SimpleNamespace(__path__=["fake"]), fake_module, Exception("nope")] + + with tracer.start_as_current_span("llamaindex.predict") as span: + apply_predict_end_safety(event, span) + instrument_llm_safety_wrappers() + + assert event.output == "[MASKED:secret]" + wrap_mock.assert_any_call( + "llama_index.core.base.llms.base", + "BaseLLM.chat", + llm_chat_wrapper, + ) + wrap_mock.assert_any_call( + "llama_index.core.base.llms.base", + "BaseLLM.complete", + llm_complete_wrapper, + ) + wrap_mock.assert_any_call( + "llama_index.core.base.llms.base", + "BaseLLM.achat", + llm_achat_wrapper, + ) + wrap_mock.assert_any_call( + "llama_index.core.base.llms.base", + "BaseLLM.acomplete", + llm_acomplete_wrapper, + ) + assert wrap_mock.call_count >= 2 + assert _resolve_masked_text("same", None) == ("same", False) + + +def test_internal_helpers_cover_kwargs_and_non_recording_paths(): + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "MEDIUM", "MASK", "PII.llamaindex", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.llamaindex", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + message = SimpleNamespace( + role=MessageRole.USER, + blocks=[SimpleNamespace(block_type="thinking", content="secret")], + ) + _, updated_kwargs = _apply_chat_prompt_safety(_FakeLLM(), (), {"messages": [message]}) + assert updated_kwargs["messages"][0].blocks[0].content == "[MASKED:secret]" + assert _apply_chat_prompt_safety(_FakeLLM(), (), {"messages": "invalid"}) == ((), {"messages": "invalid"}) + + _, updated_prompt_kwargs = _apply_completion_prompt_safety( + _FakeLLM(), (), {"prompt": "secret"} + ) + assert updated_prompt_kwargs["prompt"] == "[MASKED:secret]" + assert _apply_completion_prompt_safety(_FakeLLM(), (), {"prompt": 123}) == ((), {"prompt": 123}) + + exporter, tracer = _test_span() + with patch("opentelemetry.instrumentation.llamaindex.safety.should_send_prompts", return_value=False): + with tracer.start_as_current_span("llamaindex.completion") as span: + apply_completion_start_attributes( + LLMCompletionStartEvent( + prompt="prompt", + additional_kwargs={}, + model_dict={"llm": {"model": "test-model", "temperature": 0.1}}, + span_id="span-2", + ), + span, + ) + apply_completion_end_safety( + LLMCompletionEndEvent( + prompt="prompt", + response=CompletionResponse( + text="secret", + raw={"model": "test-model", "usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3}}, + ), + span_id="span-2", + ), + span, + ) + + spans = exporter.get_finished_spans() + assert spans[0].attributes["gen_ai.response.model"] == "test-model" + assert spans[0].attributes["gen_ai.usage.input_tokens"] == 1 + assert spans[0].attributes["gen_ai.usage.output_tokens"] == 2 + assert spans[0].attributes["llm.usage.total_tokens"] == 3 + thinking_block = SimpleNamespace(block_type="thinking", content="idea") + assert _block_text(thinking_block) == "idea" + assert _set_block_text(thinking_block, "updated") is True + assert thinking_block.content == "updated" diff --git a/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/__init__.py b/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/__init__.py index e84ac9bf7c..5496e1094f 100644 --- a/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/__init__.py +++ b/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/__init__.py @@ -13,6 +13,10 @@ ChoiceEvent, MessageEvent, ) +from opentelemetry.instrumentation.mistralai.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.mistralai.utils import ( dont_throw, should_emit_events, @@ -412,6 +416,7 @@ def _wrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, llm_request_type, name) _handle_input(span, event_logger, args, kwargs, to_wrap) response = wrapped(*args, **kwargs) @@ -422,6 +427,7 @@ def _wrap( span, event_logger, llm_request_type, response ) + _apply_completion_safety(span, response, llm_request_type, name) _handle_response(span, event_logger, llm_request_type, response) if span.is_recording(): @@ -458,6 +464,7 @@ async def _awrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, llm_request_type, name) _handle_input(span, event_logger, args, kwargs, to_wrap) if to_wrap.get("streaming"): @@ -471,6 +478,7 @@ async def _awrap( span, event_logger, llm_request_type, response ) + _apply_completion_safety(span, response, llm_request_type, name) _handle_response(span, event_logger, llm_request_type, response) if span.is_recording(): diff --git a/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/safety.py b/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/safety.py new file mode 100644 index 0000000000..4c40c0abcd --- /dev/null +++ b/packages/opentelemetry-instrumentation-mistralai/opentelemetry/instrumentation/mistralai/safety.py @@ -0,0 +1,182 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "MistralAI" + + +def _apply_prompt_safety(span, kwargs, llm_request_type, span_name): + try: + if llm_request_type != LLMRequestTypeValues.CHAT: + return kwargs + + messages = kwargs.get("messages") + if not isinstance(messages, list): + return kwargs + + mutated_kwargs = kwargs + mutated_messages = None + for index, message in enumerate(messages): + updated_content, changed = _mask_prompt_content( + span, + get_object_value(message, "content"), + span_name=span_name, + segment_index=index, + segment_role=get_object_value(message, "role"), + ) + if not changed: + continue + if mutated_messages is None: + mutated_kwargs = dict(kwargs) + mutated_messages = clone_value(messages) + mutated_kwargs["messages"] = mutated_messages + set_object_value(mutated_messages[index], "content", updated_content) + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_completion_safety(span, response, llm_request_type, span_name): + try: + if llm_request_type != LLMRequestTypeValues.CHAT: + return + + choices = get_object_value(response, "choices") or [] + for index, choice in enumerate(choices): + message = get_object_value(choice, "message") + if message is None: + continue + updated_content, changed = _mask_completion_content( + span, + get_object_value(message, "content"), + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(message, "content", updated_content) + except Exception: + return + + +def _mask_prompt_content(span, content, *, span_name, segment_index, segment_role): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_text = get_object_value(block, "text") + if not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content(span, content, *, span_name, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_text = get_object_value(block, "text") + if not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + span_name=span_name, + segment_index=segment_index, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + span_name, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index, metadata=None): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.CHAT.value, + segment_index=segment_index, + segment_role="assistant", + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-mistralai/pyproject.toml b/packages/opentelemetry-instrumentation-mistralai/pyproject.toml index 64bd327107..e6551cda75 100644 --- a/packages/opentelemetry-instrumentation-mistralai/pyproject.toml +++ b/packages/opentelemetry-instrumentation-mistralai/pyproject.toml @@ -11,6 +11,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -71,5 +72,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-mistralai/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-mistralai/tests/test_safety_hooks.py new file mode 100644 index 0000000000..2e1f076285 --- /dev/null +++ b/packages/opentelemetry-instrumentation-mistralai/tests/test_safety_hooks.py @@ -0,0 +1,146 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.mistralai.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _mask_completion_content, + _mask_prompt_content, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv_ai import LLMRequestTypeValues + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_message_content(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"messages": [SimpleNamespace(role="user", content="secret")]} + with tracer.start_as_current_span("mistralai.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, kwargs, LLMRequestTypeValues.CHAT, "mistralai.chat" + ) + + assert updated_kwargs["messages"][0].content == "[PII.chat]" + + +def test_completion_safety_masks_choice_message(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.chat]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="secret"))]) + with tracer.start_as_current_span("mistralai.chat") as span: + _apply_completion_safety(span, response, LLMRequestTypeValues.CHAT, "mistralai.chat") + + assert response.choices[0].message.content == "[SECRET.chat]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_safety_handles_non_chat_and_content_blocks(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + + with tracer.start_as_current_span("mistralai.chat") as span: + assert _apply_prompt_safety( + span, + {"messages": [SimpleNamespace(role="user", content="secret")]}, + LLMRequestTypeValues.COMPLETION, + "mistralai.chat", + )["messages"][0].content == "secret" + updated_kwargs = _apply_prompt_safety( + span, + {"messages": [SimpleNamespace(role="user", content=[SimpleNamespace(text="secret-block")])]}, + LLMRequestTypeValues.CHAT, + "mistralai.chat", + ) + assert _mask_prompt_content( + span, None, span_name="mistralai.chat", segment_index=0, segment_role="user" + ) == (None, False) + + assert updated_kwargs["messages"][0].content[0].text == "[MASKED:secret-block]" + + +def test_completion_helpers_cover_passthrough_and_block_paths(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + response = SimpleNamespace( + choices=[ + SimpleNamespace(message=SimpleNamespace(content=None)), + SimpleNamespace(message=SimpleNamespace(content=[SimpleNamespace(text="secret-block")])), + ] + ) + with tracer.start_as_current_span("mistralai.chat") as span: + _apply_completion_safety(span, response, LLMRequestTypeValues.CHAT, "mistralai.chat") + _apply_completion_safety(span, response, LLMRequestTypeValues.COMPLETION, "mistralai.chat") + assert _mask_completion_content( + span, None, span_name="mistralai.chat", segment_index=0 + ) == (None, False) + + assert response.choices[1].message.content[0].text == "[MASKED:secret-block]" + assert _resolve_masked_text("same", None) == ("same", False) + unchanged = SafetyResult(text="same", overall_action="MASK", findings=[]) + assert _resolve_masked_text("same", unchanged) == ("same", False) + assert exporter.get_finished_spans() diff --git a/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/__init__.py b/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/__init__.py index a688c92224..1526771a9e 100644 --- a/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/__init__.py +++ b/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/__init__.py @@ -14,6 +14,10 @@ emit_choice_events, emit_message_events, ) +from opentelemetry.instrumentation.ollama.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.ollama.span_utils import ( set_input_attributes, set_model_input_attributes, @@ -306,6 +310,7 @@ def _wrap( SpanAttributes.LLM_REQUEST_TYPE: llm_request_type.value, }, ) + kwargs = _apply_prompt_safety(span, kwargs, llm_request_type, name) _handle_input(span, event_logger, llm_request_type, args, kwargs) start_time = time.perf_counter() @@ -340,6 +345,7 @@ def _wrap( start_time, ) + _apply_completion_safety(span, response, llm_request_type, name) _handle_response( span, event_logger, llm_request_type, token_histogram, response ) @@ -381,6 +387,7 @@ async def _awrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, llm_request_type, name) _handle_input(span, event_logger, llm_request_type, args, kwargs) start_time = time.perf_counter() @@ -414,6 +421,7 @@ async def _awrap( start_time, ) + _apply_completion_safety(span, response, llm_request_type, name) _handle_response( span, event_logger, llm_request_type, token_histogram, response ) diff --git a/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/safety.py b/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/safety.py new file mode 100644 index 0000000000..9bfdc989c2 --- /dev/null +++ b/packages/opentelemetry-instrumentation-ollama/opentelemetry/instrumentation/ollama/safety.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) + +PROVIDER = "Ollama" + + +def _apply_prompt_safety(span, kwargs, llm_request_type, span_name): + try: + json_data = kwargs.get("json") + if not isinstance(json_data, dict): + return kwargs + + mutated_kwargs = kwargs + mutated_json = json_data + + prompt = json_data.get("prompt") + if isinstance(prompt, str): + updated_prompt, changed = _mask_prompt_text( + span, + prompt, + request_type=llm_request_type.value, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + mutated_kwargs = dict(kwargs) + mutated_json = dict(json_data) + mutated_kwargs["json"] = mutated_json + mutated_json["prompt"] = updated_prompt + + messages = json_data.get("messages") + if not isinstance(messages, list): + return mutated_kwargs + + mutated_messages = None + for index, message in enumerate(messages): + updated_content, changed = _mask_prompt_content( + span, + get_object_value(message, "content"), + request_type=llm_request_type.value, + span_name=span_name, + segment_index=index, + segment_role=get_object_value(message, "role"), + ) + if not changed: + continue + if mutated_messages is None: + if mutated_kwargs is kwargs: + mutated_kwargs = dict(kwargs) + if mutated_json is json_data: + mutated_json = dict(json_data) + mutated_kwargs["json"] = mutated_json + mutated_messages = clone_value(messages) + mutated_json["messages"] = mutated_messages + set_object_value(mutated_messages[index], "content", updated_content) + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_completion_safety(span, response, llm_request_type, span_name): + try: + if llm_request_type.value == "chat": + message = get_object_value(response, "message") + if message is None: + return + updated_content, changed = _mask_completion_content( + span, + get_object_value(message, "content"), + request_type=llm_request_type.value, + span_name=span_name, + segment_index=0, + ) + if changed: + set_object_value(message, "content", updated_content) + return + + text = get_object_value(response, "response") + if not isinstance(text, str): + return + updated_text, changed = _mask_completion_text( + span, + text, + request_type=llm_request_type.value, + span_name=span_name, + segment_index=0, + ) + if changed: + set_object_value(response, "response", updated_text) + except Exception: + return + + +def _mask_prompt_content( + span, + content, + *, + request_type, + span_name, + segment_index, + segment_role, +): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_text = get_object_value(block, "text") + if not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content(span, content, *, request_type, span_name, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + block_text = get_object_value(block, "text") + if not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + metadata={"block_index": block_index}, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + request_type, + span_name, + segment_index, + segment_role, + metadata=None, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, request_type, span_name, segment_index, metadata=None): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role="assistant", + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-ollama/pyproject.toml b/packages/opentelemetry-instrumentation-ollama/pyproject.toml index 54822e8946..a8cdeb5268 100644 --- a/packages/opentelemetry-instrumentation-ollama/pyproject.toml +++ b/packages/opentelemetry-instrumentation-ollama/pyproject.toml @@ -11,6 +11,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -71,5 +72,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-ollama/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-ollama/tests/test_safety_hooks.py new file mode 100644 index 0000000000..03eb23c23a --- /dev/null +++ b/packages/opentelemetry-instrumentation-ollama/tests/test_safety_hooks.py @@ -0,0 +1,137 @@ +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.ollama.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _mask_completion_content, + _mask_prompt_content, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv_ai import LLMRequestTypeValues + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_chat_json_messages(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"json": {"messages": [{"role": "user", "content": "secret"}]}} + with tracer.start_as_current_span("ollama.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, kwargs, LLMRequestTypeValues.CHAT, "ollama.chat" + ) + + assert updated_kwargs["json"]["messages"][0]["content"] == "[PII.chat]" + + +def test_completion_safety_masks_chat_message(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.chat]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = {"message": {"content": "secret", "role": "assistant"}} + with tracer.start_as_current_span("ollama.chat") as span: + _apply_completion_safety(span, response, LLMRequestTypeValues.CHAT, "ollama.chat") + + assert response["message"]["content"] == "[SECRET.chat]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_and_completion_cover_prompt_blocks_and_non_chat_response(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + kwargs = { + "json": { + "prompt": "prompt-secret", + "messages": [{"role": "user", "content": [{"text": "block-secret"}]}], + } + } + response = {"response": "completion-secret"} + with tracer.start_as_current_span("ollama.generate") as span: + updated_kwargs = _apply_prompt_safety( + span, kwargs, LLMRequestTypeValues.CHAT, "ollama.generate" + ) + _apply_completion_safety( + span, response, LLMRequestTypeValues.COMPLETION, "ollama.generate" + ) + assert _mask_prompt_content( + span, + None, + request_type="chat", + span_name="ollama.generate", + segment_index=0, + segment_role="user", + ) == (None, False) + assert _mask_completion_content( + span, + None, + request_type="completion", + span_name="ollama.generate", + segment_index=0, + ) == (None, False) + + assert updated_kwargs["json"]["prompt"] == "[MASKED:prompt-secret]" + assert updated_kwargs["json"]["messages"][0]["content"][0]["text"] == "[MASKED:block-secret]" + assert response["response"] == "[MASKED:completion-secret]" + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-openai/pyproject.toml b/packages/opentelemetry-instrumentation-openai/pyproject.toml index c8e40064e1..8ae02fa5a2 100644 --- a/packages/opentelemetry-instrumentation-openai/pyproject.toml +++ b/packages/opentelemetry-instrumentation-openai/pyproject.toml @@ -72,6 +72,11 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] diff --git a/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/__init__.py b/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/__init__.py index 106c39fd3a..066138d269 100644 --- a/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/__init__.py +++ b/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/__init__.py @@ -13,6 +13,10 @@ emit_event, ) from opentelemetry.instrumentation.replicate.event_models import MessageEvent +from opentelemetry.instrumentation.replicate.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.replicate.span_utils import ( set_input_attributes, set_model_input_attributes, @@ -131,6 +135,7 @@ def _wrap( }, ) + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) _handle_request(span, event_logger, args, kwargs) response = wrapped(*args, **kwargs) @@ -139,6 +144,7 @@ def _wrap( if is_streaming_response(response): return _build_from_streaming_response(span, event_logger, response) else: + response = _apply_completion_safety(span, response, name) _handle_response(span, event_logger, response) span.end() diff --git a/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/safety.py b/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/safety.py new file mode 100644 index 0000000000..9bcdf49b82 --- /dev/null +++ b/packages/opentelemetry-instrumentation-replicate/opentelemetry/instrumentation/replicate/safety.py @@ -0,0 +1,195 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "Replicate" +_PROMPT_KEY_MARKERS = ("prompt", "text", "input", "query", "instruction") + + +def _apply_prompt_safety(span, args, kwargs, span_name): + try: + model_input = kwargs.get("input") or (args[1] if len(args) > 1 else None) + if not isinstance(model_input, dict): + return args, kwargs + + updated_input, changed = _mask_prompt_value( + span, + model_input, + span_name=span_name, + segment_index=0, + ) + if not changed: + return args, kwargs + + if "input" in kwargs: + mutated_kwargs = dict(kwargs) + mutated_kwargs["input"] = updated_input + return args, mutated_kwargs + + mutated_args = list(args) + mutated_args[1] = updated_input + return tuple(mutated_args), kwargs + except Exception: + return args, kwargs + + +def _apply_completion_safety(span, response, span_name): + try: + if isinstance(response, list): + for index, item in enumerate(response): + if not isinstance(item, str): + continue + updated_item, changed = _mask_completion_text( + span, + item, + span_name=span_name, + segment_index=index, + ) + if changed: + response[index] = updated_item + return response + + if isinstance(response, str): + updated_response, changed = _mask_completion_text( + span, + response, + span_name=span_name, + segment_index=0, + ) + return updated_response if changed else response + + output = get_object_value(response, "output") + if isinstance(output, list): + updated_output = list(output) + changed_output = False + for index, item in enumerate(output): + if not isinstance(item, str): + continue + updated_item, changed = _mask_completion_text( + span, + item, + span_name=span_name, + segment_index=index, + ) + if changed: + updated_output[index] = updated_item + changed_output = True + if changed_output: + set_object_value(response, "output", updated_output) + return response + + if not isinstance(output, str): + return response + updated_output, changed = _mask_completion_text( + span, + output, + span_name=span_name, + segment_index=0, + ) + if changed: + set_object_value(response, "output", updated_output) + return response + except Exception: + return response + + +def _mask_prompt_value(span, value, *, span_name, segment_index, key=None): + if isinstance(value, str): + if not _should_mask_prompt_key(key): + return value, False + return _mask_prompt_text( + span, + value, + span_name=span_name, + segment_index=segment_index, + segment_role="user", + ) + + if isinstance(value, dict): + updated_value = value + for child_key, child_value in value.items(): + updated_child, changed = _mask_prompt_value( + span, + child_value, + span_name=span_name, + segment_index=segment_index, + key=child_key, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[child_key] = updated_child + return updated_value, updated_value is not value + + if not isinstance(value, list): + return value, False + + updated_value = value + for index, item in enumerate(value): + updated_item, changed = _mask_prompt_value( + span, + item, + span_name=span_name, + segment_index=index if _should_mask_prompt_key(key) else segment_index, + key=key, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[index] = updated_item + + return updated_value, updated_value is not value + + +def _should_mask_prompt_key(key): + if not isinstance(key, str): + return False + normalized = key.lower() + return any(marker in normalized for marker in _PROMPT_KEY_MARKERS) + + +def _mask_prompt_text(span, text, *, span_name, segment_index, segment_role): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-replicate/pyproject.toml b/packages/opentelemetry-instrumentation-replicate/pyproject.toml index f46a2b4c24..dc4149971c 100644 --- a/packages/opentelemetry-instrumentation-replicate/pyproject.toml +++ b/packages/opentelemetry-instrumentation-replicate/pyproject.toml @@ -10,6 +10,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -69,5 +70,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-replicate/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-replicate/tests/test_safety_hooks.py new file mode 100644 index 0000000000..7eef406945 --- /dev/null +++ b/packages/opentelemetry-instrumentation-replicate/tests/test_safety_hooks.py @@ -0,0 +1,192 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.replicate.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.instrumentation.replicate import _wrap +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_input_prompt(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + with tracer.start_as_current_span("replicate.run") as span: + _, updated_kwargs = _apply_prompt_safety( + span, (), {"input": {"prompt": "secret"}}, "replicate.run" + ) + + assert updated_kwargs["input"]["prompt"] == "[PII.prompt]" + + +def test_completion_safety_masks_prediction_output(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(output="secret") + with tracer.start_as_current_span("replicate.predictions.create") as span: + _apply_completion_safety(span, response, "replicate.predictions.create") + + assert response.output == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_safety_supports_args_path_and_passthrough(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + with tracer.start_as_current_span("replicate.run") as span: + updated_args, unchanged_kwargs = _apply_prompt_safety( + span, + ("model", {"prompt": "secret"}), + {}, + "replicate.run", + ) + _, nested_kwargs = _apply_prompt_safety( + span, + (), + { + "input": { + "query": "secret", + "image": "https://example.invalid/cat.png", + "nested": {"negative_prompt": "secret"}, + } + }, + "replicate.run", + ) + passthrough_args, passthrough_kwargs = _apply_prompt_safety( + span, + (), + {"input": {"prompt": None}}, + "replicate.run", + ) + + assert updated_args[1]["prompt"] == "[PII.prompt]" + assert unchanged_kwargs == {} + assert nested_kwargs["input"]["query"] == "[PII.prompt]" + assert nested_kwargs["input"]["nested"]["negative_prompt"] == "[PII.prompt]" + assert nested_kwargs["input"]["image"] == "https://example.invalid/cat.png" + assert passthrough_args == () + assert passthrough_kwargs == {"input": {"prompt": None}} + + +def test_completion_safety_covers_list_and_string_branches(): + _, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + list_response = ["secret", 123] + object_response = SimpleNamespace(output=["secret", "second"]) + with tracer.start_as_current_span("replicate.predictions.create") as span: + updated_string = _apply_completion_safety( + span, + "plain-string", + "replicate.predictions.create", + ) + _apply_completion_safety(span, list_response, "replicate.predictions.create") + _apply_completion_safety(span, object_response, "replicate.predictions.create") + + assert list_response[0] == "[MASKED:secret]" + assert object_response.output == ["[MASKED:secret]", "[MASKED:second]"] + assert updated_string == "[MASKED:plain-string]" + assert _resolve_masked_text("same", None) == ("same", False) + unchanged = SafetyResult(text="same", overall_action="MASK", findings=[]) + assert _resolve_masked_text("same", unchanged) == ("same", False) + + +def test_wrapper_masks_plain_string_response_before_attributes_are_written(): + exporter, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text="[MASKED:plain-string]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "plain-string" + else None + ) + + wrapper = _wrap(tracer, None, {"span_name": "replicate.run"}) + response = wrapper( + lambda *args, **kwargs: "plain-string", + None, + (), + {"input": {"prompt": "secret"}}, + ) + + assert response == "[MASKED:plain-string]" + span = exporter.get_finished_spans()[0] + assert span.attributes["gen_ai.prompt.0.user"] == "[PII.prompt]" + assert span.attributes["gen_ai.completion.0.content"] == "[MASKED:plain-string]" diff --git a/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/__init__.py b/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/__init__.py index 2c83f3a3c9..f3e6c85294 100644 --- a/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/__init__.py +++ b/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/__init__.py @@ -1,6 +1,7 @@ """OpenTelemetry SageMaker instrumentation""" from functools import wraps +from io import BytesIO from typing import Collection from opentelemetry import context as context_api @@ -14,6 +15,10 @@ from opentelemetry.instrumentation.sagemaker.reusable_streaming_body import ( ReusableStreamingBody, ) +from opentelemetry.instrumentation.sagemaker.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.sagemaker.span_utils import ( set_call_request_attributes, set_call_response_attributes, @@ -96,6 +101,7 @@ def with_instrumentation(*args, **kwargs): with tracer.start_as_current_span( "sagemaker.completion", kind=SpanKind.CLIENT ) as span: + kwargs = _apply_prompt_safety(span, kwargs, "sagemaker.completion") response = fn(*args, **kwargs) _handle_call(span, event_logger, kwargs, response) @@ -159,10 +165,16 @@ def _handle_call(span, event_logger, kwargs, response): set_call_span_attributes(span, kwargs, response) # Handle Response + raw_response = response["Body"].read() + raw_response, changed = _apply_completion_safety( + span, raw_response, "sagemaker.completion" + ) + if changed: + response["Body"] = ReusableStreamingBody(BytesIO(raw_response), len(raw_response)) + if should_emit_events() and event_logger is not None: emit_choice_events(response, event_logger) else: - raw_response = response["Body"].read() set_call_response_attributes(span, raw_response) diff --git a/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/safety.py b/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/safety.py new file mode 100644 index 0000000000..f3dc792e2f --- /dev/null +++ b/packages/opentelemetry-instrumentation-sagemaker/opentelemetry/instrumentation/sagemaker/safety.py @@ -0,0 +1,193 @@ +from __future__ import annotations + +import json + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + run_completion_safety, + run_prompt_safety, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "SageMaker" +_PROMPT_KEYS = {"prompt", "inputs", "inputText", "input_text", "text"} +_COMPLETION_KEYS = {"generated_text", "text", "completion", "output", "response"} + + +def _apply_prompt_safety(span, kwargs, span_name): + body = kwargs.get("Body") + payload, as_bytes = _decode_payload(body) + if payload is None: + return kwargs + + masked_payload, changed = _mask_prompt_payload( + span, + payload, + span_name=span_name, + segment_index=0, + ) + if not changed: + return kwargs + + mutated_kwargs = dict(kwargs) + mutated_kwargs["Body"] = _encode_payload(masked_payload, as_bytes) + return mutated_kwargs + + +def _apply_completion_safety(span, raw_response, span_name): + payload, as_bytes = _decode_payload(raw_response) + if payload is None: + return raw_response, False + + masked_payload, changed = _mask_completion_payload( + span, + payload, + span_name=span_name, + segment_index=0, + ) + if not changed: + return raw_response, False + + return _encode_payload(masked_payload, as_bytes), True +def _mask_prompt_payload(span, value, *, span_name, segment_index): + if isinstance(value, dict): + updated = value + for key, item in value.items(): + if key in _PROMPT_KEYS and isinstance(item, str): + updated_item, changed = _mask_prompt_text( + span, + item, + span_name=span_name, + segment_index=segment_index, + segment_role="user", + ) + else: + updated_item, changed = _mask_prompt_payload( + span, + item, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated is value: + updated = dict(value) + updated[key] = updated_item + return updated, updated is not value + + if isinstance(value, list): + updated = value + for index, item in enumerate(value): + updated_item, changed = _mask_prompt_payload( + span, + item, + span_name=span_name, + segment_index=index, + ) + if not changed: + continue + if updated is value: + updated = list(value) + updated[index] = updated_item + return updated, updated is not value + + return value, False + + +def _mask_completion_payload(span, value, *, span_name, segment_index): + if isinstance(value, dict): + updated = value + for key, item in value.items(): + if key in _COMPLETION_KEYS and isinstance(item, str): + updated_item, changed = _mask_completion_text( + span, + item, + span_name=span_name, + segment_index=segment_index, + ) + else: + updated_item, changed = _mask_completion_payload( + span, + item, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated is value: + updated = dict(value) + updated[key] = updated_item + return updated, updated is not value + + if isinstance(value, list): + updated = value + for index, item in enumerate(value): + updated_item, changed = _mask_completion_payload( + span, + item, + span_name=span_name, + segment_index=index, + ) + if not changed: + continue + if updated is value: + updated = list(value) + updated[index] = updated_item + return updated, updated is not value + + return value, False + + +def _mask_prompt_text(span, text, *, span_name, segment_index, segment_role): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _decode_payload(raw_value): + try: + if isinstance(raw_value, bytes): + return json.loads(raw_value.decode("utf-8")), True + if isinstance(raw_value, str): + return json.loads(raw_value), False + except Exception: + return None, False + return None, False + + +def _encode_payload(payload, as_bytes): + encoded = json.dumps(payload).encode("utf-8") + if as_bytes: + return encoded + return encoded.decode("utf-8") + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-sagemaker/pyproject.toml b/packages/opentelemetry-instrumentation-sagemaker/pyproject.toml index 904c8cdd5c..7f5130f378 100644 --- a/packages/opentelemetry-instrumentation-sagemaker/pyproject.toml +++ b/packages/opentelemetry-instrumentation-sagemaker/pyproject.toml @@ -10,6 +10,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -67,5 +68,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-sagemaker/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-sagemaker/tests/test_safety_hooks.py new file mode 100644 index 0000000000..5b3eb5f00c --- /dev/null +++ b/packages/opentelemetry-instrumentation-sagemaker/tests/test_safety_hooks.py @@ -0,0 +1,130 @@ +import json +from io import BytesIO + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.sagemaker.reusable_streaming_body import ( + ReusableStreamingBody, +) +from opentelemetry.instrumentation.sagemaker.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _decode_payload, + _encode_payload, + _mask_completion_payload, + _mask_prompt_payload, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_request_body(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"Body": json.dumps({"inputs": "secret"})} + with tracer.start_as_current_span("sagemaker.completion") as span: + updated_kwargs = _apply_prompt_safety(span, kwargs, "sagemaker.completion") + + assert json.loads(updated_kwargs["Body"])["inputs"] == "[PII.prompt]" + + +def test_completion_safety_masks_response_body(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + raw_response = json.dumps([{"generated_text": "secret"}]) + with tracer.start_as_current_span("sagemaker.completion") as span: + updated_response, changed = _apply_completion_safety( + span, raw_response, "sagemaker.completion" + ) + + assert changed is True + assert json.loads(updated_response)[0]["generated_text"] == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_payload_helpers_cover_nested_and_invalid_inputs(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + payload = {"items": [{"inputs": "prompt-a"}]} + completion_payload = {"items": [{"generated_text": "completion-a"}]} + with tracer.start_as_current_span("sagemaker.completion") as span: + masked_prompt, prompt_changed = _mask_prompt_payload( + span, payload, span_name="sagemaker.completion", segment_index=0 + ) + masked_completion, completion_changed = _mask_completion_payload( + span, completion_payload, span_name="sagemaker.completion", segment_index=0 + ) + + assert prompt_changed is True + assert masked_prompt["items"][0]["inputs"] == "[MASKED:prompt-a]" + assert completion_changed is True + assert masked_completion["items"][0]["generated_text"] == "[MASKED:completion-a]" + assert _decode_payload("not-json") == (None, False) + assert _decode_payload(123) == (None, False) + assert _encode_payload({"text": "x"}, False) == '{"text": "x"}' + assert _encode_payload({"text": "x"}, True) == b'{"text": "x"}' + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/__init__.py b/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/__init__.py index 985effb250..ff9bab42af 100644 --- a/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/__init__.py +++ b/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/__init__.py @@ -11,6 +11,10 @@ emit_completion_event, emit_prompt_events, ) +from opentelemetry.instrumentation.together.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.together.span_utils import ( set_completion_attributes, set_model_completion_attributes, @@ -120,11 +124,13 @@ def _wrap( SpanAttributes.LLM_REQUEST_TYPE: llm_request_type.value, }, ) + kwargs = _apply_prompt_safety(span, kwargs, llm_request_type, name) _handle_input(span, event_logger, llm_request_type, kwargs) response = wrapped(*args, **kwargs) if response: + _apply_completion_safety(span, response, llm_request_type, name) _handle_response(span, event_logger, llm_request_type, response) if span.is_recording(): span.set_status(Status(StatusCode.OK)) diff --git a/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/safety.py b/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/safety.py new file mode 100644 index 0000000000..ffd65633e7 --- /dev/null +++ b/packages/opentelemetry-instrumentation-together/opentelemetry/instrumentation/together/safety.py @@ -0,0 +1,249 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) + +PROVIDER = "TogetherAI" + + +def _apply_prompt_safety(span, kwargs, llm_request_type, span_name): + try: + mutated_kwargs = kwargs + + prompt = kwargs.get("prompt") + if isinstance(prompt, str): + updated_prompt, changed = _mask_prompt_text( + span, + prompt, + request_type=llm_request_type.value, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + mutated_kwargs = dict(kwargs) + mutated_kwargs["prompt"] = updated_prompt + + messages = kwargs.get("messages") + if not isinstance(messages, list): + return mutated_kwargs + + mutated_messages = None + for index, message in enumerate(messages): + updated_content, changed = _mask_prompt_content( + span, + get_object_value(message, "content"), + request_type=llm_request_type.value, + span_name=span_name, + segment_index=index, + segment_role=get_object_value(message, "role"), + ) + if not changed: + continue + if mutated_messages is None: + if mutated_kwargs is kwargs: + mutated_kwargs = dict(kwargs) + mutated_messages = clone_value(messages) + mutated_kwargs["messages"] = mutated_messages + set_object_value(mutated_messages[index], "content", updated_content) + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_completion_safety(span, response, llm_request_type, span_name): + try: + choices = get_object_value(response, "choices") or [] + for index, choice in enumerate(choices): + if llm_request_type.value == "chat": + message = get_object_value(choice, "message") + if message is None: + continue + updated_content, changed = _mask_completion_content( + span, + get_object_value(message, "content"), + request_type=llm_request_type.value, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(message, "content", updated_content) + continue + + text = get_object_value(choice, "text") + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + span, + text, + request_type=llm_request_type.value, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(choice, "text", updated_text) + except Exception: + return + + +def _mask_prompt_content( + span, + content, + *, + request_type, + span_name, + segment_index, + segment_role, +): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_prompt_text( + span, + block, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content(span, content, *, request_type, span_name, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_completion_text( + span, + block, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + request_type, + span_name, + segment_index, + segment_role, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, request_type, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-together/pyproject.toml b/packages/opentelemetry-instrumentation-together/pyproject.toml index ba86188f19..6175210634 100644 --- a/packages/opentelemetry-instrumentation-together/pyproject.toml +++ b/packages/opentelemetry-instrumentation-together/pyproject.toml @@ -12,6 +12,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -72,5 +73,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-together/pytest.ini b/packages/opentelemetry-instrumentation-together/pytest.ini index 40880458c7..55e2214826 100644 --- a/packages/opentelemetry-instrumentation-together/pytest.ini +++ b/packages/opentelemetry-instrumentation-together/pytest.ini @@ -1,2 +1,4 @@ [pytest] asyncio_mode=auto +markers = + safety: safety-focused tests diff --git a/packages/opentelemetry-instrumentation-together/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-together/tests/test_safety_hooks.py new file mode 100644 index 0000000000..ac46139743 --- /dev/null +++ b/packages/opentelemetry-instrumentation-together/tests/test_safety_hooks.py @@ -0,0 +1,192 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.together.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv_ai import LLMRequestTypeValues + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_chat_message(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"messages": [{"role": "user", "content": "secret"}]} + with tracer.start_as_current_span("together.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, kwargs, LLMRequestTypeValues.CHAT, "together.chat" + ) + + assert updated_kwargs["messages"][0]["content"] == "[PII.chat]" + + +def test_completion_safety_masks_completion_choice(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(choices=[SimpleNamespace(text="secret")]) + with tracer.start_as_current_span("together.completion") as span: + _apply_completion_safety( + span, response, LLMRequestTypeValues.COMPLETION, "together.completion" + ) + + assert response.choices[0].text == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_and_completion_cover_prompt_and_chat_branches(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + with tracer.start_as_current_span("together.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, + {"prompt": "prompt-secret", "messages": [{"role": "user", "content": "message-secret"}]}, + LLMRequestTypeValues.CHAT, + "together.chat", + ) + response = SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="message-secret"))]) + _apply_completion_safety(span, response, LLMRequestTypeValues.CHAT, "together.chat") + unchanged = _apply_prompt_safety( + span, {"messages": "invalid"}, LLMRequestTypeValues.CHAT, "together.chat" + ) + + assert updated_kwargs["prompt"] == "[MASKED:prompt-secret]" + assert updated_kwargs["messages"][0]["content"] == "[MASKED:message-secret]" + assert response.choices[0].message.content == "[MASKED:message-secret]" + assert unchanged["messages"] == "invalid" + assert _resolve_masked_text("same", None) == ("same", False) + assert _resolve_masked_text("same", SafetyResult(text="same", overall_action="MASK", findings=[])) == ( + "same", + False, + ) + + +def test_prompt_and_completion_mask_mixed_content_blocks(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + kwargs = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "prompt-secret"}, + {"type": "image_url", "text": "leave-image-alone"}, + ], + } + ] + } + response = SimpleNamespace( + choices=[ + SimpleNamespace( + message=SimpleNamespace( + content=[ + SimpleNamespace(type="text", text="completion-secret"), + SimpleNamespace(type="image_url", text="leave-image-alone"), + ] + ) + ) + ] + ) + + with tracer.start_as_current_span("together.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, + kwargs, + LLMRequestTypeValues.CHAT, + "together.chat", + ) + _apply_completion_safety( + span, + response, + LLMRequestTypeValues.CHAT, + "together.chat", + ) + + assert updated_kwargs["messages"][0]["content"][0]["text"] == "[MASKED:prompt-secret]" + assert updated_kwargs["messages"][0]["content"][1]["text"] == "leave-image-alone" + assert response.choices[0].message.content[0].text == "[MASKED:completion-secret]" + assert response.choices[0].message.content[1].text == "leave-image-alone" diff --git a/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/safety.py b/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/safety.py new file mode 100644 index 0000000000..72ce4cb53b --- /dev/null +++ b/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/safety.py @@ -0,0 +1,198 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "Transformers" +_PROMPT_KWARGS = ("text_inputs", "args") + + +def _apply_prompt_safety(span, args, kwargs, span_name): + try: + prompts, source = _get_prompt_input(args, kwargs) + if source is None: + return args, kwargs + + updated_prompts, changed = _mask_prompt_value( + span, + prompts, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if not changed: + return args, kwargs + + if source == "args": + return (updated_prompts, *args[1:]), kwargs + + mutated_kwargs = dict(kwargs) + mutated_kwargs[source] = updated_prompts + return args, mutated_kwargs + except Exception: + return args, kwargs + + +def _apply_completion_safety(span, response, span_name): + try: + _mask_completion_value(span, response, span_name=span_name, segment_index=0) + except Exception: + return + + +def _get_prompt_input(args, kwargs): + if args: + return args[0], "args" + for key in _PROMPT_KWARGS: + if key in kwargs: + return kwargs.get(key), key + return None, None + + +def _mask_prompt_value(span, value, *, span_name, segment_index, segment_role): + if isinstance(value, str): + return _mask_prompt_text( + span, + value, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if isinstance(value, dict): + content = get_object_value(value, "content") + if not isinstance(content, str): + return value, False + + updated_content, changed = _mask_prompt_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + segment_role=get_object_value(value, "role") or segment_role, + ) + if not changed: + return value, False + updated_value = clone_value(value) + set_object_value(updated_value, "content", updated_content) + return updated_value, True + + if not isinstance(value, list): + return value, False + + updated_value = value + for index, item in enumerate(value): + updated_item, changed = _mask_prompt_value( + span, + item, + span_name=span_name, + segment_index=index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[index] = updated_item + + return updated_value, updated_value is not value + + +def _mask_completion_value(span, value, *, span_name, segment_index): + if isinstance(value, str): + return _mask_completion_text( + span, + value, + span_name=span_name, + segment_index=segment_index, + ) + + if isinstance(value, dict): + changed_any = False + generated_text = get_object_value(value, "generated_text") + if isinstance(generated_text, (str, list, dict)): + updated_generated_text, changed = _mask_completion_value( + span, + generated_text, + span_name=span_name, + segment_index=segment_index, + ) + if changed: + set_object_value(value, "generated_text", updated_generated_text) + changed_any = True + + content = get_object_value(value, "content") + if isinstance(content, str): + updated_content, changed = _mask_completion_text( + span, + content, + span_name=span_name, + segment_index=segment_index, + ) + if changed: + set_object_value(value, "content", updated_content) + changed_any = True + + return value, changed_any + + if not isinstance(value, list): + return value, False + + changed_any = False + for index, item in enumerate(value): + updated_item, changed = _mask_completion_value( + span, + item, + span_name=span_name, + segment_index=index, + ) + if not changed: + continue + value[index] = updated_item + changed_any = True + + return value, changed_any + + +def _mask_prompt_text(span, text, *, span_name, segment_index, segment_role): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/text_generation_pipeline_wrapper.py b/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/text_generation_pipeline_wrapper.py index e67fd3d86c..4a30e29ee1 100644 --- a/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/text_generation_pipeline_wrapper.py +++ b/packages/opentelemetry-instrumentation-transformers/opentelemetry/instrumentation/transformers/text_generation_pipeline_wrapper.py @@ -10,6 +10,10 @@ set_model_input_attributes, set_response_attributes, ) +from opentelemetry.instrumentation.transformers.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.transformers.utils import ( dont_throw, should_emit_events, @@ -59,11 +63,13 @@ def text_generation_pipeline_wrapper( name = to_wrap.get("span_name") with tracer.start_as_current_span(name) as span: + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) _handle_input(span, event_logger, instance, args, kwargs) response = wrapped(*args, **kwargs) if response: + _apply_completion_safety(span, response, name) _handle_response(span, event_logger, response) if span.is_recording(): span.set_status(Status(StatusCode.OK)) diff --git a/packages/opentelemetry-instrumentation-transformers/pyproject.toml b/packages/opentelemetry-instrumentation-transformers/pyproject.toml index 00f1b82d76..851a28578a 100644 --- a/packages/opentelemetry-instrumentation-transformers/pyproject.toml +++ b/packages/opentelemetry-instrumentation-transformers/pyproject.toml @@ -12,6 +12,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -65,5 +66,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-transformers/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-transformers/tests/test_safety_hooks.py new file mode 100644 index 0000000000..fcc75350ca --- /dev/null +++ b/packages/opentelemetry-instrumentation-transformers/tests/test_safety_hooks.py @@ -0,0 +1,131 @@ +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.transformers.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_text_generation_input(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + with tracer.start_as_current_span("transformers_text_generation_pipeline.call") as span: + updated_args, _ = _apply_prompt_safety( + span, ("secret",), {}, "transformers_text_generation_pipeline.call" + ) + + assert updated_args[0] == "[PII.prompt]" + + +def test_completion_safety_masks_generated_text(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = [{"generated_text": "secret"}] + with tracer.start_as_current_span("transformers_text_generation_pipeline.call") as span: + _apply_completion_safety(span, response, "transformers_text_generation_pipeline.call") + + assert response[0]["generated_text"] == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_and_completion_cover_keyword_and_nested_batch_paths(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + with tracer.start_as_current_span("transformers_text_generation_pipeline.call") as span: + _, updated_kwargs = _apply_prompt_safety( + span, + (), + {"text_inputs": ["prompt-a", "prompt-b", 123]}, + "transformers_text_generation_pipeline.call", + ) + unchanged_args, unchanged_kwargs = _apply_prompt_safety( + span, (), {"args": {"not": "supported"}}, "transformers_text_generation_pipeline.call" + ) + _, updated_chat_kwargs = _apply_prompt_safety( + span, + (), + {"text_inputs": [{"role": "user", "content": "chat-secret"}]}, + "transformers_text_generation_pipeline.call", + ) + response = [ + [{"generated_text": "completion-a"}], + [{"generated_text": [{"role": "assistant", "content": "chat-completion"}]}], + {"generated_text": None}, + "ignored", + ] + _apply_completion_safety(span, response, "transformers_text_generation_pipeline.call") + _apply_completion_safety(span, {"generated_text": "ignored"}, "transformers_text_generation_pipeline.call") + + assert updated_kwargs["text_inputs"][:2] == ["[MASKED:prompt-a]", "[MASKED:prompt-b]"] + assert updated_chat_kwargs["text_inputs"][0]["content"] == "[MASKED:chat-secret]" + assert unchanged_args == () + assert unchanged_kwargs == {"args": {"not": "supported"}} + assert response[0][0]["generated_text"] == "[MASKED:completion-a]" + assert response[1][0]["generated_text"][0]["content"] == "[MASKED:chat-completion]" + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/__init__.py b/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/__init__.py index 5d96fad8b3..c6cf82d295 100644 --- a/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/__init__.py +++ b/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/__init__.py @@ -20,6 +20,10 @@ set_model_response_attributes, set_response_attributes, ) +from opentelemetry.instrumentation.vertexai.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.vertexai.utils import dont_throw, should_emit_events from opentelemetry.instrumentation.vertexai.version import __version__ from opentelemetry.semconv._incubating.attributes import ( @@ -243,6 +247,7 @@ async def _awrap(tracer, event_logger, to_wrap, wrapped, instance, args, kwargs) }, ) + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) await _handle_request(span, event_logger, args, kwargs, llm_model) response = await wrapped(*args, **kwargs) @@ -257,6 +262,7 @@ async def _awrap(tracer, event_logger, to_wrap, wrapped, instance, args, kwargs) span, event_logger, response, llm_model ) else: + _apply_completion_safety(span, response, name) _handle_response(span, event_logger, response, llm_model) span.end() @@ -292,6 +298,7 @@ def _wrap(tracer, event_logger, to_wrap, wrapped, instance, args, kwargs): }, ) + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) # Use sync version for non-async wrapper to avoid image processing for now set_model_input_attributes(span, kwargs, llm_model) if should_emit_events(): @@ -311,6 +318,7 @@ def _wrap(tracer, event_logger, to_wrap, wrapped, instance, args, kwargs): span, event_logger, response, llm_model ) else: + _apply_completion_safety(span, response, name) _handle_response(span, event_logger, response, llm_model) span.end() diff --git a/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/safety.py b/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/safety.py new file mode 100644 index 0000000000..c6fc0ec9c3 --- /dev/null +++ b/packages/opentelemetry-instrumentation-vertexai/opentelemetry/instrumentation/vertexai/safety.py @@ -0,0 +1,200 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "VertexAI" + + +def _apply_prompt_safety(span, args, kwargs, span_name): + try: + updated_args = args + updated_kwargs = kwargs + + if args: + masked_arg, changed = _mask_prompt_value( + span, + args[0], + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + updated_args = (masked_arg, *args[1:]) + + if "contents" in kwargs: + masked_contents, changed = _mask_prompt_value( + span, + kwargs.get("contents"), + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + updated_kwargs = dict(kwargs) + updated_kwargs["contents"] = masked_contents + + return updated_args, updated_kwargs + except Exception: + return args, kwargs + + +def _apply_completion_safety(span, response, span_name): + try: + text = get_object_value(response, "text") + if isinstance(text, str): + updated_text, changed = _mask_completion_text( + span, + text, + span_name=span_name, + segment_index=0, + ) + if changed: + set_object_value(response, "text", updated_text) + + candidates = get_object_value(response, "candidates") + if not isinstance(candidates, list): + return + + for candidate_index, candidate in enumerate(candidates): + candidate_text = get_object_value(candidate, "text") + if isinstance(candidate_text, str): + updated_candidate_text, changed = _mask_completion_text( + span, + candidate_text, + span_name=span_name, + segment_index=candidate_index, + ) + if changed: + set_object_value(candidate, "text", updated_candidate_text) + + content = get_object_value(candidate, "content") + parts = get_object_value(content, "parts") + if not isinstance(parts, list): + continue + for part_index, part in enumerate(parts): + part_text = get_object_value(part, "text") + if not isinstance(part_text, str): + continue + updated_part_text, changed = _mask_completion_text( + span, + part_text, + span_name=span_name, + segment_index=candidate_index, + metadata={"part_index": part_index}, + ) + if changed: + set_object_value(part, "text", updated_part_text) + except Exception: + return + + +def _mask_prompt_value(span, value, *, span_name, segment_index, segment_role): + if isinstance(value, str): + return _mask_prompt_text( + span, + value, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if isinstance(value, list): + updated_value = value + for index, item in enumerate(value): + masked_item, changed = _mask_prompt_value( + span, + item, + span_name=span_name, + segment_index=index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_value is value: + updated_value = clone_value(value) + updated_value[index] = masked_item + return updated_value, updated_value is not value + + parts = get_object_value(value, "parts") + if isinstance(parts, list): + updated_parts = parts + updated_value = value + for index, part in enumerate(parts): + masked_part, changed = _mask_prompt_value( + span, + part, + span_name=span_name, + segment_index=index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_parts is parts: + updated_parts = clone_value(parts) + updated_value = clone_value(value) + set_object_value(updated_value, "parts", updated_parts) + updated_parts[index] = masked_part + return updated_value, updated_parts is not parts + + text = get_object_value(value, "text") + if not isinstance(text, str): + return value, False + + updated_text, changed = _mask_prompt_text( + span, + text, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + return value, False + updated_value = clone_value(value) + set_object_value(updated_value, "text", updated_text) + return updated_value, True + + +def _mask_prompt_text(span, text, *, span_name, segment_index, segment_role): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index, metadata=None): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + metadata=metadata, + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-vertexai/pyproject.toml b/packages/opentelemetry-instrumentation-vertexai/pyproject.toml index d779be2708..27e00e3f02 100644 --- a/packages/opentelemetry-instrumentation-vertexai/pyproject.toml +++ b/packages/opentelemetry-instrumentation-vertexai/pyproject.toml @@ -13,6 +13,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -61,6 +62,9 @@ show_missing = true [tool.pytest.ini_options] asyncio_mode = "auto" +markers = [ + "safety: safety-focused tests", +] [tool.ruff] line-length = 120 @@ -78,3 +82,6 @@ select = ["E", "F", "W"] [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-vertexai/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-vertexai/tests/test_safety_hooks.py new file mode 100644 index 0000000000..fccb1d9c2b --- /dev/null +++ b/packages/opentelemetry-instrumentation-vertexai/tests/test_safety_hooks.py @@ -0,0 +1,125 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.vertexai.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _mask_prompt_value, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_positional_text_arg(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + with tracer.start_as_current_span("vertexai.generate_content") as span: + updated_args, _ = _apply_prompt_safety( + span, ("secret",), {}, "vertexai.generate_content" + ) + + assert updated_args[0] == "[PII.prompt]" + + +def test_completion_safety_masks_response_text(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(text="secret", candidates=[SimpleNamespace(text="secret")]) + with tracer.start_as_current_span("vertexai.generate_content") as span: + _apply_completion_safety(span, response, "vertexai.generate_content") + + assert response.text == "[SECRET.output]" + assert response.candidates[0].text == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 2 + + +def test_prompt_and_completion_cover_contents_parts_and_candidate_parts(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + contents = [SimpleNamespace(parts=[SimpleNamespace(text="prompt-part")])] + response = SimpleNamespace( + text=None, + candidates=[ + SimpleNamespace( + text=None, + content=SimpleNamespace(parts=[SimpleNamespace(text="completion-part")]), + ) + ], + ) + with tracer.start_as_current_span("vertexai.generate_content") as span: + _, updated_kwargs = _apply_prompt_safety( + span, (), {"contents": contents}, "vertexai.generate_content" + ) + _apply_completion_safety(span, response, "vertexai.generate_content") + assert _mask_prompt_value( + span, None, span_name="vertexai.generate_content", segment_index=0, segment_role="user" + ) == (None, False) + + assert updated_kwargs["contents"][0].parts[0].text == "[MASKED:prompt-part]" + assert response.candidates[0].content.parts[0].text == "[MASKED:completion-part]" + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/__init__.py b/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/__init__.py index 8590081c4a..e5909c424a 100644 --- a/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/__init__.py +++ b/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/__init__.py @@ -18,6 +18,10 @@ emit_event, ) from opentelemetry.instrumentation.watsonx.event_models import ChoiceEvent, MessageEvent +from opentelemetry.instrumentation.watsonx.safety import ( + _apply_completion_safety, + _apply_prompt_safety, +) from opentelemetry.instrumentation.watsonx.utils import ( dont_throw, should_emit_events, @@ -578,6 +582,7 @@ def _wrap( }, ) + args, kwargs = _apply_prompt_safety(span, args, kwargs, name) _handle_input(span, event_logger, name, instance, args, kwargs) if "generate" in name: @@ -620,6 +625,7 @@ def _wrap( ) else: duration = end_time - start_time + _apply_completion_safety(span, response, name) _handle_response( span, event_logger, diff --git a/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/safety.py b/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/safety.py new file mode 100644 index 0000000000..8102c18c73 --- /dev/null +++ b/packages/opentelemetry-instrumentation-watsonx/opentelemetry/instrumentation/watsonx/safety.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + run_completion_safety, + run_prompt_safety, +) +from opentelemetry.semconv_ai import LLMRequestTypeValues + +PROVIDER = "Watsonx" + + +def _apply_prompt_safety(span, args, kwargs, span_name): + try: + prompt = kwargs.get("prompt") + if prompt is None and args: + first_arg = args[0] + if isinstance(first_arg, (str, list)): + prompt = first_arg + + if isinstance(prompt, str): + updated_prompt, changed = _mask_prompt_text( + span, + prompt, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if not changed: + return args, kwargs + if "prompt" in kwargs: + mutated_kwargs = dict(kwargs) + mutated_kwargs["prompt"] = updated_prompt + return args, mutated_kwargs + return (updated_prompt, *args[1:]), kwargs + + if not isinstance(prompt, list): + return args, kwargs + + updated_prompt = None + for index, item in enumerate(prompt): + if not isinstance(item, str): + continue + masked_item, changed = _mask_prompt_text( + span, + item, + span_name=span_name, + segment_index=index, + segment_role="user", + ) + if not changed: + continue + if updated_prompt is None: + updated_prompt = list(prompt) + updated_prompt[index] = masked_item + + if updated_prompt is None: + return args, kwargs + if "prompt" in kwargs: + mutated_kwargs = dict(kwargs) + mutated_kwargs["prompt"] = updated_prompt + return args, mutated_kwargs + return (updated_prompt, *args[1:]), kwargs + except Exception: + return args, kwargs + + +def _apply_completion_safety(span, responses, span_name): + try: + if isinstance(responses, list): + for index, response in enumerate(responses): + _apply_completion_to_response(span, response, span_name, index) + return + _apply_completion_to_response(span, responses, span_name, 0) + except Exception: + return + + +def _apply_completion_to_response(span, response, span_name, segment_index): + if not isinstance(response, dict): + return + results = response.get("results") + if not isinstance(results, list) or not results: + return + + updated_results = None + for result_index, result in enumerate(results): + if not isinstance(result, dict): + continue + generated_text = result.get("generated_text") + if not isinstance(generated_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + generated_text, + span_name=span_name, + segment_index=result_index, + ) + if not changed: + continue + if updated_results is None: + updated_results = list(results) + updated_result = dict(updated_results[result_index]) + updated_result["generated_text"] = updated_text + updated_results[result_index] = updated_result + + if updated_results is not None: + response["results"] = updated_results + + +def _mask_prompt_text(span, text, *, span_name, segment_index, segment_role): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=LLMRequestTypeValues.COMPLETION.value, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-watsonx/pyproject.toml b/packages/opentelemetry-instrumentation-watsonx/pyproject.toml index 2302b4bf25..732088873a 100644 --- a/packages/opentelemetry-instrumentation-watsonx/pyproject.toml +++ b/packages/opentelemetry-instrumentation-watsonx/pyproject.toml @@ -10,6 +10,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.13,<0.5.0", "opentelemetry-semantic-conventions>=0.59b0", @@ -62,5 +63,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "requests>=2.32.5", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-watsonx/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-watsonx/tests/test_safety_hooks.py new file mode 100644 index 0000000000..a4655ba77d --- /dev/null +++ b/packages/opentelemetry-instrumentation-watsonx/tests/test_safety_hooks.py @@ -0,0 +1,116 @@ +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.watsonx.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_prompt_arg(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.prompt]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + with tracer.start_as_current_span("watsonx.generate") as span: + updated_args, _ = _apply_prompt_safety(span, ("secret",), {}, "watsonx.generate") + + assert updated_args[0] == "[PII.prompt]" + + +def test_completion_safety_masks_generated_text(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = {"results": [{"generated_text": "secret"}]} + with tracer.start_as_current_span("watsonx.generate") as span: + _apply_completion_safety(span, response, "watsonx.generate") + + assert response["results"][0]["generated_text"] == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_and_completion_cover_list_paths(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.prompt", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + with tracer.start_as_current_span("watsonx.generate") as span: + _, updated_kwargs = _apply_prompt_safety( + span, (), {"prompt": ["a", "b", 1]}, "watsonx.generate" + ) + updated_args, _ = _apply_prompt_safety( + span, (["x", "y"],), {}, "watsonx.generate" + ) + responses = [ + {"results": [{"generated_text": "one"}, {"generated_text": "one-b"}]}, + {"results": [{"generated_text": "two"}]}, + ] + _apply_completion_safety(span, responses, "watsonx.generate") + + assert updated_kwargs["prompt"][:2] == ["[MASKED:a]", "[MASKED:b]"] + assert updated_args[0][:2] == ["[MASKED:x]", "[MASKED:y]"] + assert responses[0]["results"][0]["generated_text"] == "[MASKED:one]" + assert responses[0]["results"][1]["generated_text"] == "[MASKED:one-b]" + assert responses[1]["results"][0]["generated_text"] == "[MASKED:two]" + assert _resolve_masked_text("same", None) == ("same", False) diff --git a/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/__init__.py b/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/__init__.py index c087236b3b..af41f47541 100644 --- a/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/__init__.py +++ b/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/__init__.py @@ -26,6 +26,8 @@ from opentelemetry.instrumentation.writer.config import Config from opentelemetry.instrumentation.writer.event_emitter import ( emit_choice_events, emit_message_events) +from opentelemetry.instrumentation.writer.safety import ( + _apply_completion_safety, _apply_prompt_safety) from opentelemetry.instrumentation.writer.span_utils import ( set_input_attributes, set_model_input_attributes, set_model_response_attributes, set_response_attributes) @@ -361,6 +363,7 @@ def _wrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, request_type, name) _handle_input(span, kwargs, event_logger) start_time = time.time() @@ -413,6 +416,9 @@ def _wrap( attributes=response_attributes(response, to_wrap.get("method")), ) + _apply_completion_safety( + span, response, request_type, name + ) _handle_response( span, response, token_histogram, event_logger, to_wrap.get("method") ) @@ -462,6 +468,7 @@ async def _awrap( }, ) + kwargs = _apply_prompt_safety(span, kwargs, request_type, name) _handle_input(span, kwargs, event_logger) start_time = time.time() @@ -514,6 +521,9 @@ async def _awrap( attributes=response_attributes(response, to_wrap.get("method")), ) + _apply_completion_safety( + span, response, request_type, name + ) _handle_response( span, response, token_histogram, event_logger, to_wrap.get("method") ) diff --git a/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/safety.py b/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/safety.py new file mode 100644 index 0000000000..5facd80c37 --- /dev/null +++ b/packages/opentelemetry-instrumentation-writer/opentelemetry/instrumentation/writer/safety.py @@ -0,0 +1,247 @@ +from __future__ import annotations + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyDecision, + SafetyLocation, + clone_value, + get_object_value, + run_completion_safety, + run_prompt_safety, + set_object_value, +) + +PROVIDER = "Writer" + + +def _apply_prompt_safety(span, kwargs, request_type, span_name): + try: + mutated_kwargs = kwargs + + prompt = kwargs.get("prompt") + if isinstance(prompt, str): + updated_prompt, changed = _mask_prompt_text( + span, + prompt, + request_type=request_type.value, + span_name=span_name, + segment_index=0, + segment_role="user", + ) + if changed: + mutated_kwargs = dict(kwargs) + mutated_kwargs["prompt"] = updated_prompt + + messages = kwargs.get("messages") + if not isinstance(messages, list): + return mutated_kwargs + + mutated_messages = None + for index, message in enumerate(messages): + updated_content, changed = _mask_prompt_content( + span, + get_object_value(message, "content"), + request_type=request_type.value, + span_name=span_name, + segment_index=index, + segment_role=get_object_value(message, "role"), + ) + if not changed: + continue + if mutated_messages is None: + if mutated_kwargs is kwargs: + mutated_kwargs = dict(kwargs) + mutated_messages = clone_value(messages) + mutated_kwargs["messages"] = mutated_messages + set_object_value(mutated_messages[index], "content", updated_content) + + return mutated_kwargs + except Exception: + return kwargs + + +def _apply_completion_safety(span, response, request_type, span_name): + try: + choices = get_object_value(response, "choices") or [] + for index, choice in enumerate(choices): + message = get_object_value(choice, "message") + if message is not None: + updated_content, changed = _mask_completion_content( + span, + get_object_value(message, "content"), + request_type=request_type.value, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(message, "content", updated_content) + continue + + text = get_object_value(choice, "text") + if not isinstance(text, str): + continue + updated_text, changed = _mask_completion_text( + span, + text, + request_type=request_type.value, + span_name=span_name, + segment_index=index, + ) + if changed: + set_object_value(choice, "text", updated_text) + except Exception: + return + + +def _mask_prompt_content( + span, + content, + *, + request_type, + span_name, + segment_index, + segment_role, +): + if isinstance(content, str): + return _mask_prompt_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_prompt_text( + span, + block, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_prompt_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + segment_role=segment_role, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_completion_content(span, content, *, request_type, span_name, segment_index): + if isinstance(content, str): + return _mask_completion_text( + span, + content, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + + if not isinstance(content, list): + return content, False + + updated_content = content + for block_index, block in enumerate(content): + if isinstance(block, str): + updated_text, changed = _mask_completion_text( + span, + block, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + updated_content[block_index] = updated_text + continue + + block_type = get_object_value(block, "type") + block_text = get_object_value(block, "text") + if block_type not in (None, "text") or not isinstance(block_text, str): + continue + updated_text, changed = _mask_completion_text( + span, + block_text, + request_type=request_type, + span_name=span_name, + segment_index=segment_index, + ) + if not changed: + continue + if updated_content is content: + updated_content = clone_value(content) + set_object_value(updated_content[block_index], "text", updated_text) + + return updated_content, updated_content is not content + + +def _mask_prompt_text( + span, + text, + *, + request_type, + span_name, + segment_index, + segment_role, +): + result = run_prompt_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.PROMPT, + request_type=request_type, + segment_index=segment_index, + segment_role=segment_role, + ) + return _resolve_masked_text(text, result) + + +def _mask_completion_text(span, text, *, request_type, span_name, segment_index): + result = run_completion_safety( + span=span, + provider=PROVIDER, + span_name=span_name, + text=text, + location=SafetyLocation.COMPLETION, + request_type=request_type, + segment_index=segment_index, + segment_role="assistant", + ) + return _resolve_masked_text(text, result) + + +def _resolve_masked_text(original_text, result): + if result is None or result.overall_action != SafetyDecision.MASK.value: + return original_text, False + if result.text == original_text: + return original_text, False + return result.text, True diff --git a/packages/opentelemetry-instrumentation-writer/pyproject.toml b/packages/opentelemetry-instrumentation-writer/pyproject.toml index 6bc4384848..86be6d32d3 100644 --- a/packages/opentelemetry-instrumentation-writer/pyproject.toml +++ b/packages/opentelemetry-instrumentation-writer/pyproject.toml @@ -11,6 +11,7 @@ readme = "README.md" requires-python = ">=3.10,<4" dependencies = [ "opentelemetry-api>=1.38.0,<2", + "opentelemetry-instrumentation-fortifyroot", "opentelemetry-instrumentation>=0.59b0", "opentelemetry-semantic-conventions-ai>=0.4.11", "opentelemetry-semantic-conventions>=0.59b0", @@ -72,5 +73,13 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "W"] +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.uv] constraint-dependencies = ["urllib3>=2.6.3", "pyarrow>=18.1.0", "pip>=25.3"] + +[tool.uv.sources] +opentelemetry-instrumentation-fortifyroot = { path = "../opentelemetry-instrumentation-fortifyroot", editable = true } diff --git a/packages/opentelemetry-instrumentation-writer/tests/test_safety_hooks.py b/packages/opentelemetry-instrumentation-writer/tests/test_safety_hooks.py new file mode 100644 index 0000000000..486362817e --- /dev/null +++ b/packages/opentelemetry-instrumentation-writer/tests/test_safety_hooks.py @@ -0,0 +1,192 @@ +from types import SimpleNamespace + +import pytest + +from opentelemetry.instrumentation.fortifyroot import ( + SafetyFinding, + SafetyLocation, + SafetyResult, + clear_safety_handlers, + register_completion_safety_handler, + register_prompt_safety_handler, +) +from opentelemetry.instrumentation.writer.safety import ( + _apply_completion_safety, + _apply_prompt_safety, + _resolve_masked_text, +) +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.semconv_ai import LLMRequestTypeValues + +pytestmark = pytest.mark.safety + + +def setup_function(): + clear_safety_handlers() + + +def teardown_function(): + clear_safety_handlers() + + +def _test_span(): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + return exporter, tracer + + +def test_prompt_safety_masks_chat_message(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text="[PII.chat]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT and context.text == "secret" + else None + ) + + kwargs = {"messages": [{"role": "user", "content": "secret"}]} + with tracer.start_as_current_span("writerai.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, kwargs, LLMRequestTypeValues.CHAT, "writerai.chat" + ) + + assert updated_kwargs["messages"][0]["content"] == "[PII.chat]" + + +def test_completion_safety_masks_choice_text(): + exporter, tracer = _test_span() + register_completion_safety_handler( + lambda context: SafetyResult( + text="[SECRET.output]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION and context.text == "secret" + else None + ) + + response = SimpleNamespace(choices=[SimpleNamespace(text="secret")]) + with tracer.start_as_current_span("writerai.completions") as span: + _apply_completion_safety( + span, response, LLMRequestTypeValues.COMPLETION, "writerai.completions" + ) + + assert response.choices[0].text == "[SECRET.output]" + assert len(exporter.get_finished_spans()[0].events) == 1 + + +def test_prompt_and_completion_cover_prompt_and_message_paths(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + with tracer.start_as_current_span("writerai.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, + {"prompt": "prompt-secret", "messages": [{"role": "user", "content": "message-secret"}]}, + LLMRequestTypeValues.CHAT, + "writerai.chat", + ) + response = SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="message-secret"))]) + _apply_completion_safety(span, response, LLMRequestTypeValues.CHAT, "writerai.chat") + unchanged = _apply_prompt_safety( + span, {"messages": "invalid"}, LLMRequestTypeValues.CHAT, "writerai.chat" + ) + + assert updated_kwargs["prompt"] == "[MASKED:prompt-secret]" + assert updated_kwargs["messages"][0]["content"] == "[MASKED:message-secret]" + assert response.choices[0].message.content == "[MASKED:message-secret]" + assert unchanged["messages"] == "invalid" + assert _resolve_masked_text("same", None) == ("same", False) + assert _resolve_masked_text("same", SafetyResult(text="same", overall_action="MASK", findings=[])) == ( + "same", + False, + ) + + +def test_prompt_and_completion_mask_mixed_content_blocks(): + _, tracer = _test_span() + register_prompt_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("PII", "HIGH", "MASK", "PII.chat", 0, len(context.text))], + ) + if context.location == SafetyLocation.PROMPT + else None + ) + register_completion_safety_handler( + lambda context: SafetyResult( + text=f"[MASKED:{context.text}]", + overall_action="MASK", + findings=[SafetyFinding("SECRET", "HIGH", "MASK", "SECRET.output", 0, len(context.text))], + ) + if context.location == SafetyLocation.COMPLETION + else None + ) + + kwargs = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "prompt-secret"}, + {"type": "image_url", "text": "leave-image-alone"}, + ], + } + ] + } + response = SimpleNamespace( + choices=[ + SimpleNamespace( + message=SimpleNamespace( + content=[ + SimpleNamespace(type="text", text="completion-secret"), + SimpleNamespace(type="image_url", text="leave-image-alone"), + ] + ) + ) + ] + ) + + with tracer.start_as_current_span("writerai.chat") as span: + updated_kwargs = _apply_prompt_safety( + span, + kwargs, + LLMRequestTypeValues.CHAT, + "writerai.chat", + ) + _apply_completion_safety( + span, + response, + LLMRequestTypeValues.CHAT, + "writerai.chat", + ) + + assert updated_kwargs["messages"][0]["content"][0]["text"] == "[MASKED:prompt-secret]" + assert updated_kwargs["messages"][0]["content"][1]["text"] == "leave-image-alone" + assert response.choices[0].message.content[0].text == "[MASKED:completion-secret]" + assert response.choices[0].message.content[1].text == "leave-image-alone" diff --git a/packages/opentelemetry-semantic-conventions-ai/opentelemetry/semconv_ai/__init__.py b/packages/opentelemetry-semantic-conventions-ai/opentelemetry/semconv_ai/__init__.py index 6d8356b77c..0cae486211 100644 --- a/packages/opentelemetry-semantic-conventions-ai/opentelemetry/semconv_ai/__init__.py +++ b/packages/opentelemetry-semantic-conventions-ai/opentelemetry/semconv_ai/__init__.py @@ -88,6 +88,9 @@ class SpanAttributes: LLM_CONTENT_COMPLETION_CHUNK = "llm.content.completion.chunk" LLM_REQUEST_REASONING_EFFORT = "llm.request.reasoning_effort" LLM_USAGE_REASONING_TOKENS = "llm.usage.reasoning_tokens" + # Backward-compatible aliases still used across instrumentation packages. + LLM_PROMPTS = "gen_ai.prompt" + LLM_COMPLETIONS = "gen_ai.completion" # OpenAI LLM_OPENAI_RESPONSE_SYSTEM_FINGERPRINT = "gen_ai.openai.system_fingerprint" diff --git a/packages/traceloop-sdk/pyproject.toml b/packages/traceloop-sdk/pyproject.toml index f094f0fcb5..373fa5b56b 100644 --- a/packages/traceloop-sdk/pyproject.toml +++ b/packages/traceloop-sdk/pyproject.toml @@ -39,6 +39,7 @@ dependencies = [ "opentelemetry-instrumentation-transformers", "opentelemetry-instrumentation-together", "opentelemetry-instrumentation-llamaindex", + "opentelemetry-instrumentation-litellm", "opentelemetry-instrumentation-milvus", "opentelemetry-instrumentation-haystack", "opentelemetry-instrumentation-bedrock", @@ -106,6 +107,11 @@ source = ["traceloop/sdk"] exclude_lines = ["if TYPE_CHECKING:"] show_missing = true +[tool.pytest.ini_options] +markers = [ + "safety: safety-focused tests", +] + [tool.mypy] python_version = "3.10" warn_return_any = true @@ -169,6 +175,7 @@ opentelemetry-instrumentation-haystack = { path = "../opentelemetry-instrumentat opentelemetry-instrumentation-lancedb = { path = "../opentelemetry-instrumentation-lancedb", editable = true } opentelemetry-instrumentation-langchain = { path = "../opentelemetry-instrumentation-langchain", editable = true } opentelemetry-instrumentation-llamaindex = { path = "../opentelemetry-instrumentation-llamaindex", editable = true } +opentelemetry-instrumentation-litellm = { path = "../opentelemetry-instrumentation-litellm", editable = true } opentelemetry-instrumentation-marqo = { path = "../opentelemetry-instrumentation-marqo", editable = true } opentelemetry-instrumentation-mcp = { path = "../opentelemetry-instrumentation-mcp", editable = true } opentelemetry-instrumentation-milvus = { path = "../opentelemetry-instrumentation-milvus", editable = true } diff --git a/packages/traceloop-sdk/tests/test_litellm_instrumentation.py b/packages/traceloop-sdk/tests/test_litellm_instrumentation.py new file mode 100644 index 0000000000..ab913de822 --- /dev/null +++ b/packages/traceloop-sdk/tests/test_litellm_instrumentation.py @@ -0,0 +1,42 @@ +from unittest.mock import patch + +import pytest + + +@pytest.mark.safety +def test_init_litellm_instrumentor_instruments_when_package_is_present(): + from traceloop.sdk.tracing.tracing import init_litellm_instrumentor + + with patch("traceloop.sdk.tracing.tracing.is_package_installed", return_value=True), patch( + "opentelemetry.instrumentation.litellm.LiteLLMInstrumentor" + ) as instrumentor_cls: + instrumentor = instrumentor_cls.return_value + instrumentor.is_instrumented_by_opentelemetry = False + + assert init_litellm_instrumentor() is True + instrumentor.instrument.assert_called_once_with() + + +@pytest.mark.safety +def test_init_litellm_instrumentor_returns_false_when_package_is_missing(): + from traceloop.sdk.tracing.tracing import init_litellm_instrumentor + + with patch("traceloop.sdk.tracing.tracing.is_package_installed", return_value=False): + assert init_litellm_instrumentor() is False + + +@pytest.mark.safety +def test_init_instrumentations_dispatches_litellm(): + from traceloop.sdk.instruments import Instruments + from traceloop.sdk.tracing.tracing import init_instrumentations + + with patch( + "traceloop.sdk.tracing.tracing.init_litellm_instrumentor", return_value=True + ) as init_mock: + init_instrumentations( + should_enrich_metrics=False, + base64_image_uploader=lambda *args, **kwargs: "", + instruments={Instruments.LITELLM}, + ) + + init_mock.assert_called_once_with() diff --git a/packages/traceloop-sdk/tests/test_sdk_initialization.py b/packages/traceloop-sdk/tests/test_sdk_initialization.py index b518b92ee0..d1475938da 100644 --- a/packages/traceloop-sdk/tests/test_sdk_initialization.py +++ b/packages/traceloop-sdk/tests/test_sdk_initialization.py @@ -215,3 +215,4 @@ def test_get_default_span_processor(): assert isinstance(processor, BatchSpanProcessor) assert hasattr(processor, "_traceloop_processor") assert getattr(processor, "_traceloop_processor") is True + diff --git a/packages/traceloop-sdk/traceloop/sdk/instruments.py b/packages/traceloop-sdk/traceloop/sdk/instruments.py index 6d1d27bd66..775045f2c6 100644 --- a/packages/traceloop-sdk/traceloop/sdk/instruments.py +++ b/packages/traceloop-sdk/traceloop/sdk/instruments.py @@ -16,6 +16,7 @@ class Instruments(Enum): LANCEDB = "lancedb" LANGCHAIN = "langchain" LLAMA_INDEX = "llama_index" + LITELLM = "litellm" MARQO = "marqo" MCP = "mcp" MILVUS = "milvus" diff --git a/packages/traceloop-sdk/traceloop/sdk/tracing/tracing.py b/packages/traceloop-sdk/traceloop/sdk/tracing/tracing.py index cbd44e28d4..2d659882d4 100644 --- a/packages/traceloop-sdk/traceloop/sdk/tracing/tracing.py +++ b/packages/traceloop-sdk/traceloop/sdk/tracing/tracing.py @@ -533,6 +533,9 @@ def init_instrumentations( elif instrument == Instruments.LLAMA_INDEX: if init_llama_index_instrumentor(): instrument_set = True + elif instrument == Instruments.LITELLM: + if init_litellm_instrumentor(): + instrument_set = True elif instrument == Instruments.MARQO: if init_marqo_instrumentor(): instrument_set = True @@ -844,6 +847,20 @@ def init_llama_index_instrumentor(): return False +def init_litellm_instrumentor(): + try: + if is_package_installed("litellm"): + from opentelemetry.instrumentation.litellm import LiteLLMInstrumentor + + instrumentor = LiteLLMInstrumentor() + if not instrumentor.is_instrumented_by_opentelemetry: + instrumentor.instrument() + return True + except Exception as e: + logging.error(f"Error initializing LiteLLM instrumentor: {e}") + return False + + def init_milvus_instrumentor(): try: if is_package_installed("pymilvus"):