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
46 changes: 46 additions & 0 deletions .github/workflows/safety-pr.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
name: Safety PR Tests

on:
pull_request:

permissions:
contents: read

concurrency:
group: safety-pr-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true

jobs:
safety-tests:
name: Safety Tests
runs-on: ubuntu-latest
timeout-minutes: 45

steps:
- name: Check out code
uses: actions/checkout@v4
with:
fetch-depth: 0
ref: ${{ github.event.pull_request.head.sha }}

- name: Set up Python 3.11
uses: actions/setup-python@v5
with:
python-version: "3.11"
cache: "pip"

- name: Install Poetry
run: python -m pip install --upgrade pip poetry

- name: Run safety test suite
env:
HAYSTACK_TELEMETRY_ENABLED: "False"
run: bash ./scripts/run-tests.sh --safety

- name: Upload safety test reports
if: always()
uses: actions/upload-artifact@v4
with:
name: safety-test-reports
path: reports/test-run/
if-no-files-found: ignore
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,10 @@
emit_input_events,
emit_response_events,
)
from opentelemetry.instrumentation.anthropic.safety import (
_apply_completion_safety,
_apply_prompt_safety,
)
from opentelemetry.instrumentation.anthropic.span_utils import (
aset_input_attributes,
set_response_attributes,
Expand Down Expand Up @@ -540,8 +544,8 @@ def _wrap(
},
)

kwargs = _apply_prompt_safety(span, kwargs, name)
_handle_input(span, event_logger, kwargs)

start_time = time.time()
try:
response = wrapped(*args, **kwargs)
Expand Down Expand Up @@ -611,6 +615,7 @@ def _wrap(
attributes=metric_attributes,
)

_apply_completion_safety(span, response, name)
_handle_response(span, event_logger, response)
if span.is_recording():
_set_token_usage(
Expand Down Expand Up @@ -663,8 +668,8 @@ async def _awrap(
SpanAttributes.LLM_REQUEST_TYPE: LLMRequestTypeValues.COMPLETION.value,
},
)
kwargs = _apply_prompt_safety(span, kwargs, name)
await _ahandle_input(span, event_logger, kwargs)

start_time = time.time()
try:
response = await wrapped(*args, **kwargs)
Expand Down Expand Up @@ -735,6 +740,7 @@ async def _awrap(
attributes=metric_attributes,
)

_apply_completion_safety(span, response, name)
await _ahandle_response(span, event_logger, response)

if span.is_recording():
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,233 @@
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 = "Anthropic"


def _apply_prompt_safety(span, kwargs, span_name: str):
try:
request_type = _request_type(span_name)
mutated_kwargs = kwargs

prompt = kwargs.get("prompt")
if isinstance(prompt, str):
updated_prompt, changed = _mask_prompt_text(
span,
prompt,
span_name=span_name,
request_type=request_type,
segment_index=0,
segment_role="user",
)
if changed:
mutated_kwargs = dict(kwargs)
mutated_kwargs["prompt"] = updated_prompt

system = kwargs.get("system")
updated_system, system_changed = _mask_prompt_content(
span,
system,
span_name=span_name,
request_type=request_type,
segment_index=0,
segment_role="system",
)
if system_changed:
if mutated_kwargs is kwargs:
mutated_kwargs = dict(kwargs)
mutated_kwargs["system"] = updated_system

messages = kwargs.get("messages")
if not isinstance(messages, list):
return mutated_kwargs

mutated_messages = None
for index, message in enumerate(messages):
role = get_object_value(message, "role")
content = get_object_value(message, "content")
updated_content, changed = _mask_prompt_content(
span,
content,
span_name=span_name,
request_type=request_type,
segment_index=index,
segment_role=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 _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 = content
for block_index, block in enumerate(content):
if get_object_value(block, "type") != "text":
continue
text = get_object_value(block, "text")
if not isinstance(text, str):
continue
updated_text, changed = _mask_prompt_text(
span,
text,
span_name=span_name,
request_type=request_type,
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 _apply_completion_safety(span, response, span_name: str):
try:
request_type = _request_type(span_name)

completion = get_object_value(response, "completion")
if isinstance(completion, str):
updated_completion, changed = _mask_completion_text(
span,
completion,
span_name=span_name,
request_type=request_type,
segment_index=0,
segment_role="assistant",
)
if changed:
set_object_value(response, "completion", updated_completion)

content = get_object_value(response, "content")
if not isinstance(content, list):
return

for index, block in enumerate(content):
block_type = get_object_value(block, "type")
text_key = None
role = "assistant"
if block_type == "text":
text_key = "text"
elif block_type == "thinking":
text_key = "thinking"
role = "thinking"
if text_key is None:
continue
text = get_object_value(block, text_key)
if not isinstance(text, str):
continue
updated_text, changed = _mask_completion_text(
span,
text,
span_name=span_name,
request_type=request_type,
segment_index=index,
segment_role=role,
)
if changed:
set_object_value(block, text_key, updated_text)
except Exception:
return


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,
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 _request_type(span_name: str) -> str:
if span_name.endswith("completion"):
return LLMRequestTypeValues.COMPLETION.value
return LLMRequestTypeValues.CHAT.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
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.14,<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 }
Loading
Loading