Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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():
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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 }
Original file line number Diff line number Diff line change
@@ -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)
Loading
Loading