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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"

[project]
name = "simpleaudit"
version = "0.3.1"
version = "0.3.0"
description = "Lightweight AI Safety Auditing Framework"
readme = "README.md"
license = {text = "MIT"}
Expand Down
14 changes: 14 additions & 0 deletions simpleaudit/targets/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,7 @@ async def send(
documents=documents,
response_format=response_format,
params=params,
context=context,
max_retries=self.max_retries,
retry_backoff=self.retry_backoff,
)
Expand All @@ -104,16 +105,29 @@ async def _client_send(
documents: Optional[List[Union[str, Dict[str, Any]]]],
response_format: Optional[Dict[str, Any]],
params: Optional[Dict[str, Any]],
context: Optional[TargetContext] = None,
max_retries: int,
retry_backoff: float,
) -> TargetResponse:
"""Delegate to the shared AnyLLM call path.

Imported lazily so that ``simpleaudit.targets`` does not create an import
cycle with ``model_auditor`` at package import time.

When the per-turn context carries W3C trace headers (e.g. the engine's
scenario traceparent), they are merged into ``extra_headers`` so the
target's outgoing call propagates the trace context. This is what lets an
instrumented target export spans under the same trace id the engine's
TraceCorrelation records. User-supplied extra_headers are preserved.
"""
from ..model_auditor import ModelAuditor

if context is not None and context.trace_headers:
headers = dict(params.get("extra_headers", {}) if params else {})
headers.update(context.trace_headers)
params = dict(params or {})
params["extra_headers"] = headers

content, input_tokens, output_tokens = await ModelAuditor._call_async(
client,
model,
Expand Down
69 changes: 69 additions & 0 deletions tests/test_targets.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,40 @@ async def test_model_target_client_mode():
assert client.calls[0]["model"] == "m"


@pytest.mark.asyncio
async def test_model_target_client_mode_forwards_trace_headers():
"""trace_headers on the per-turn context reach the wire as extra_headers."""
client = _FakeClient()
t = ModelTarget(client=client, model="m")
ctx = TargetContext(trace_headers={"traceparent": "00-11223344556677889900112233445566-abcdef0123456789-01"})
await t.send(user="hi", context=ctx)
assert client.calls[0]["extra_headers"] == {
"traceparent": "00-11223344556677889900112233445566-abcdef0123456789-01"
}


@pytest.mark.asyncio
async def test_model_target_trace_headers_merge_with_user_params():
"""User-supplied extra_headers and the framework's trace headers coexist."""
client = _FakeClient()
t = ModelTarget(client=client, model="m")
ctx = TargetContext(trace_headers={"traceparent": "00-11223344556677889900112233445566-abcdef0123456789-01"})
await t.send(user="hi", params={"extra_headers": {"X-Custom": "1"}}, context=ctx)
assert client.calls[0]["extra_headers"] == {
"X-Custom": "1",
"traceparent": "00-11223344556677889900112233445566-abcdef0123456789-01",
}


@pytest.mark.asyncio
async def test_model_target_no_context_no_extra_headers():
"""Without a context the call is byte-identical to the legacy path."""
client = _FakeClient()
t = ModelTarget(client=client, model="m")
await t.send(user="hi")
assert "extra_headers" not in client.calls[0]


@pytest.mark.asyncio
async def test_model_target_transport_mode():
async def transport(**kw):
Expand Down Expand Up @@ -324,3 +358,38 @@ async def send(self, *, user, history=None, context=None, **kw):
assert parts[3] == "01"
# The correlation recorded at least one turn -> trace link.
assert correlation.all_trace_ids()


def test_engine_traceparent_reaches_wire_on_default_target():
"""With the default ModelTarget (no override), the scenario's traceparent
is forwarded on the target's outgoing acompletion as extra_headers."""
import asyncio

from simpleaudit.tracing.context import TraceCorrelation
from tests.fakes import fixed_probe_auditor, fixed_severity_judge, make_auditor

client = _FakeClient()
auditor = make_auditor(
target=client,
judge=fixed_severity_judge("pass"),
auditor=fixed_probe_auditor("probe"),
max_turns=2,
show_progress=False,
)

correlation = TraceCorrelation(audit_run_id="audit_test")
scenarios = [{"name": "Wire", "description": "wire traceparent test"}]
asyncio.run(
auditor.run_async(
scenarios=scenarios,
max_turns=2,
audit_run_id="audit_test",
trace_correlation=correlation,
)
)

assert client.calls
last = client.calls[-1]
tp = last["extra_headers"]["traceparent"]
# Same trace id the correlation recorded, not a fresh unrelated one.
assert tp.split("-")[1] in correlation.all_trace_ids()
20 changes: 20 additions & 0 deletions tests/test_tracing.py
Original file line number Diff line number Diff line change
Expand Up @@ -422,6 +422,26 @@ def test_parse_otlp_json_string_body():
assert len(spans) == 1


def test_parse_otlp_json_preserves_genai_content_attributes():
payload = _otlp_http_payload()
payload["resourceSpans"][0]["scopeSpans"][0]["spans"][0]["attributes"].extend(
[
{"key": "gen_ai.input.messages", "value": {"stringValue": "user prompt"}},
{"key": "gen_ai.output.messages", "value": {"stringValue": "assistant output"}},
{"key": "gen_ai.tool.call.arguments", "value": {"stringValue": '{"x":1}'}},
{
"key": "gen_ai.retrieval.documents",
"value": {"arrayValue": {"values": [{"stringValue": "document text"}]}},
},
]
)
span = parse_otlp_json(payload)[0]
assert span["attributes"]["gen_ai.input.messages"] == "user prompt"
assert span["attributes"]["gen_ai.output.messages"] == "assistant output"
assert span["attributes"]["gen_ai.tool.call.arguments"] == '{"x":1}'
assert span["attributes"]["gen_ai.retrieval.documents"] == ["document text"]


@pytest.mark.asyncio
async def test_otlp_receiver_ingests():
receiver = OTLPTraceReceiver()
Expand Down
Loading