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 @@ -3,6 +3,7 @@
from opentelemetry.instrumentation.fortifyroot import (
SafetyDecision,
SafetyLocation,
build_safety_metadata,
clone_value,
get_object_value,
run_completion_safety,
Expand All @@ -17,6 +18,7 @@
def _apply_prompt_safety(span, kwargs, span_name: str):
try:
request_type = _request_type(span_name)
request_model = kwargs.get("model")
mutated_kwargs = kwargs

prompt = kwargs.get("prompt")
Expand All @@ -28,6 +30,7 @@ def _apply_prompt_safety(span, kwargs, span_name: str):
request_type=request_type,
segment_index=0,
segment_role="user",
request_model=request_model,
)
if changed:
mutated_kwargs = dict(kwargs)
Expand All @@ -41,6 +44,7 @@ def _apply_prompt_safety(span, kwargs, span_name: str):
request_type=request_type,
segment_index=0,
segment_role="system",
request_model=request_model,
)
if system_changed:
if mutated_kwargs is kwargs:
Expand All @@ -62,6 +66,7 @@ def _apply_prompt_safety(span, kwargs, span_name: str):
request_type=request_type,
segment_index=index,
segment_role=role,
request_model=request_model,
)
if not changed:
continue
Expand All @@ -85,6 +90,7 @@ def _mask_prompt_content(
request_type,
segment_index,
segment_role,
request_model=None,
):
if isinstance(content, str):
return _mask_prompt_text(
Expand All @@ -94,6 +100,7 @@ def _mask_prompt_content(
request_type=request_type,
segment_index=segment_index,
segment_role=segment_role,
request_model=request_model,
)

if not isinstance(content, list):
Expand All @@ -114,6 +121,7 @@ def _mask_prompt_content(
segment_index=segment_index,
segment_role=segment_role,
metadata={"block_index": block_index},
request_model=request_model,
)
if not changed:
continue
Expand All @@ -127,6 +135,7 @@ def _mask_prompt_content(
def _apply_completion_safety(span, response, span_name: str):
try:
request_type = _request_type(span_name)
response_model = get_object_value(response, "model")

completion = get_object_value(response, "completion")
if isinstance(completion, str):
Expand All @@ -137,6 +146,7 @@ def _apply_completion_safety(span, response, span_name: str):
request_type=request_type,
segment_index=0,
segment_role="assistant",
response_model=response_model,
)
if changed:
set_object_value(response, "completion", updated_completion)
Expand Down Expand Up @@ -166,6 +176,7 @@ def _apply_completion_safety(span, response, span_name: str):
request_type=request_type,
segment_index=index,
segment_role=role,
response_model=response_model,
)
if changed:
set_object_value(block, text_key, updated_text)
Expand All @@ -182,7 +193,13 @@ def _mask_prompt_text(
segment_index,
segment_role,
metadata=None,
request_model=None,
):
metadata = build_safety_metadata(
metadata,
provider=PROVIDER,
request_model=request_model,
)
result = run_prompt_safety(
span=span,
provider=PROVIDER,
Expand All @@ -205,7 +222,14 @@ def _mask_completion_text(
request_type,
segment_index,
segment_role,
metadata=None,
response_model=None,
):
metadata = build_safety_metadata(
metadata,
provider=PROVIDER,
response_model=response_model,
)
result = run_completion_safety(
span=span,
provider=PROVIDER,
Expand All @@ -215,6 +239,7 @@ def _mask_completion_text(
request_type=request_type,
segment_index=segment_index,
segment_role=segment_role,
metadata=metadata,
)
return _resolve_masked_text(text, result)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,21 @@ def test_apply_prompt_safety_masks_prompt_system_and_messages(monkeypatch):
assert updated["messages"][0]["content"][0]["text"] == "masked:msg-secret"


def test_apply_prompt_safety_passes_model_metadata(monkeypatch):
contexts = []

def _prompt(**kwargs):
contexts.append(kwargs)
return None

monkeypatch.setattr(safety, "run_prompt_safety", _prompt)
kwargs = {"model": "claude-3-5-sonnet", "messages": [{"role": "user", "content": "secret"}]}
safety._apply_prompt_safety(None, kwargs, "anthropic.chat")

assert contexts[0]["metadata"]["gen_ai.system"] == "Anthropic"
assert contexts[0]["metadata"]["gen_ai.request.model"] == "claude-3-5-sonnet"


def test_apply_prompt_safety_returns_partial_update_when_messages_missing():
monkeypatch = pytest.MonkeyPatch()
monkeypatch.setattr(safety, "run_prompt_safety", lambda **kwargs: SafetyResult(text=f"masked:{kwargs['text']}", overall_action="MASK"))
Expand Down Expand Up @@ -80,6 +95,21 @@ def test_apply_completion_safety_masks_completion_and_content(monkeypatch):
assert response.content[1]["thinking"] == "masked:thought-secret"


def test_apply_completion_safety_passes_model_metadata(monkeypatch):
contexts = []

def _completion(**kwargs):
contexts.append(kwargs)
return None

monkeypatch.setattr(safety, "run_completion_safety", _completion)
response = SimpleNamespace(model="claude-3-5-sonnet", completion="secret")
safety._apply_completion_safety(None, response, "anthropic.chat")

assert contexts[0]["metadata"]["gen_ai.system"] == "Anthropic"
assert contexts[0]["metadata"]["gen_ai.response.model"] == "claude-3-5-sonnet"


def test_anthropic_prompt_and_completion_helpers_cover_noop_branches(monkeypatch):
monkeypatch.setattr(safety, "run_prompt_safety", lambda **kwargs: SafetyResult(text=kwargs["text"], overall_action="MASK"))
updated, changed = safety._mask_prompt_content(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -364,12 +364,17 @@ def stream_done(response_body):
@dont_throw
def _handle_call(span: Span, kwargs, response, metric_params, event_logger):
request_body = json.loads(kwargs.get("body"))
response_body = _prepare_invoke_response(span, response, _BEDROCK_INVOKE_SPAN_NAME)
(provider, model_vendor, model) = _get_vendor_model(kwargs.get("modelId"))
response_body = _prepare_invoke_response(
span,
response,
_BEDROCK_INVOKE_SPAN_NAME,
response_model=model,
)
headers = {}
if "ResponseMetadata" in response:
headers = response.get("ResponseMetadata").get("HTTPHeaders", {})

(provider, model_vendor, model) = _get_vendor_model(kwargs.get("modelId"))
metric_params.vendor = provider
metric_params.model = model
metric_params.is_stream = False
Expand Down Expand Up @@ -400,8 +405,13 @@ 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"))
_apply_converse_completion_safety(
span,
response,
_BEDROCK_CONVERSE_SPAN_NAME,
response_model=model,
)
guardrail_converse(span, response, provider, model, metric_params)

set_converse_model_span_attributes(span, provider, model, kwargs)
Expand Down
Loading
Loading