Skip to content
Open
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
10 changes: 10 additions & 0 deletions docs/api/core-protocols.md
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,16 @@ Protocols and ABCs that define RAMPART's extension points. Implement these to co
- register_default_handler_factory
- clear_default_handler_factory

## Trace Execution

::: rampart.core.trace
options:
members:
- EvaluationRecord
- TraceRun
- run_trace_async
- evaluate_final_trace_async

## Errors

::: rampart.core.errors
Expand Down
10 changes: 10 additions & 0 deletions rampart/core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,12 @@
resolve_attack_verdict,
resolve_probe_verdict,
)
from rampart.core.trace import (
EvaluationRecord,
TraceRun,
evaluate_final_trace_async,
run_trace_async,
)
from rampart.core.types import (
EvalContext,
EvalOutcome,
Expand Down Expand Up @@ -63,6 +69,7 @@
"EvalOutcome",
"EvalResult",
"EvaluationPurpose",
"EvaluationRecord",
"Evaluator",
"ExecutionEvent",
"ExecutionEventData",
Expand Down Expand Up @@ -92,11 +99,14 @@
"ToolCall",
"ToolDeclaration",
"TraceEndReason",
"TraceRun",
"Turn",
"evaluate_final_trace_async",
"evaluate_turn_async",
"execute_trials_async",
"resolve_as_attack",
"resolve_as_probe",
"resolve_attack_verdict",
"resolve_probe_verdict",
"run_trace_async",
]
229 changes: 229 additions & 0 deletions rampart/core/trace.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,229 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

"""Shared linear trace execution and final-trace evaluation helpers."""

from __future__ import annotations

from dataclasses import dataclass, field, replace
from typing import TYPE_CHECKING

from rampart.common.text import safe_str_list
from rampart.core.types import (
EvalContext,
EvalResult,
EvaluationPurpose,
ObservabilityLevel,
TraceEndReason,
Turn,
)

if TYPE_CHECKING:
from rampart.core.adapter import Session
from rampart.core.evaluator import Evaluator
from rampart.core.manifest import AppManifest
from rampart.core.prompt_driver import PromptDriver


@dataclass(frozen=True, kw_only=True, eq=False)
class EvaluationRecord:
"""One online evaluation and the exact context it judged.

Args:
evaluator: Evaluator object that produced the result. Identity is the
reuse boundary.
context: Exact raw-trace context passed to the evaluator.
result: Evaluation returned for that context.
"""

evaluator: Evaluator
context: EvalContext
result: EvalResult


@dataclass(kw_only=True)
class TraceRun:
"""A completed linear trace and its latest online evaluation.

``turns`` is the driver/report view and may carry online evidence.
``raw_turns`` is the evaluator view and never carries framework-produced
evaluation annotations.

Args:
trace_end_reason: Why the trace stopped producing turns.
observability_level: What the adapter can observe.
manifest: Agent capabilities used to create evaluator contexts.
turns: Annotated history passed to prompt drivers and results.
raw_turns: Annotation-free history passed to evaluators.
latest_online_evaluation: Most recent stop-condition evaluation.
"""

trace_end_reason: TraceEndReason
observability_level: ObservabilityLevel
manifest: AppManifest | None = None
turns: list[Turn] = field(default_factory=list[Turn])
raw_turns: list[Turn] = field(default_factory=list[Turn])
latest_online_evaluation: EvaluationRecord | None = None


def _evaluation_context(
*,
raw_turns: list[Turn],
observability_level: ObservabilityLevel,
manifest: AppManifest | None,
) -> EvalContext:
"""Build an evaluator context from a snapshot of the raw trace.

Returns:
EvalContext: Context holding a shallow snapshot of raw turns.
"""
return EvalContext(
turns=list(raw_turns),
observability_level=observability_level,
manifest=manifest,
)


def _matches_final_trace_context(*, context: EvalContext, run: TraceRun) -> bool:
"""Check that the raw trace and adapter context are unchanged.

Returns:
bool: Whether this context can supply the final-trace judgment.
"""
return (
context.observability_level is run.observability_level
and context.manifest is run.manifest
and len(context.turns) == len(run.raw_turns)
and all(
evaluated is final
for evaluated, final in zip(context.turns, run.raw_turns, strict=True)
)
)


async def run_trace_async(
*,
session: Session,
driver: PromptDriver,
max_turns: int,
observability_level: ObservabilityLevel,
stop_when: Evaluator | None = None,
manifest: AppManifest | None = None,
) -> TraceRun:
"""Drive a linear conversation with optional online stopping.

The runner does not own session lifetime or exception conversion. Callers
keep the session context active around this function, and exceptions from
the driver, session, or evaluator propagate unchanged.

Args:
session: Active agent session.
driver: Prompt source for the conversation.
max_turns: Maximum number of requests sent to the agent.
observability_level: What the adapter can observe.
stop_when: Optional evaluator checked after every response. A detected
outcome terminates the trace.
manifest: Agent capabilities exposed to evaluators.

Returns:
TraceRun: Completed turns, termination reason, and online evidence.

Raises:
ValueError: If ``max_turns`` is negative.
"""
if max_turns < 0:
msg = "max_turns must be non-negative."
raise ValueError(msg)

run = TraceRun(
trace_end_reason=TraceEndReason.MAX_TURNS_REACHED,
observability_level=observability_level,
manifest=manifest,
)

for turn_index in range(max_turns):
decision = await driver.next_prompt_async(history=list(run.turns))
if decision is None:
run.trace_end_reason = TraceEndReason.DRIVER_EXHAUSTED
return run

response = await session.send_async(decision.request)
raw_turn = Turn(
request=decision.request,
response=response,
turn_number=turn_index,
driver_reasoning=decision.reasoning,
)
run.raw_turns.append(raw_turn)

if stop_when is None:
run.turns.append(raw_turn)
continue

context = _evaluation_context(
raw_turns=run.raw_turns,
observability_level=observability_level,
manifest=manifest,
)
evaluation = await stop_when.evaluate_async(context=context)
run.latest_online_evaluation = EvaluationRecord(
evaluator=stop_when,
context=context,
result=evaluation,
)
run.turns.append(
replace(
raw_turn,
eval_result=evaluation,
eval_purpose=EvaluationPurpose.STOP_CHECK,
),
)
if evaluation.detected:
run.trace_end_reason = TraceEndReason.STOP_CONDITION_MET
return run

return run


async def evaluate_final_trace_async(
*,
evaluator: Evaluator,
run: TraceRun,
) -> EvalResult | None:
"""Evaluate the final raw trace, reusing an identical online judgment.

Args:
evaluator: Evaluator responsible for the final verdict.
run: Completed trace from :func:`run_trace_async`.

Returns:
EvalResult | None: Final evaluation, or None when no turns exist.

Call this before leaving any active session or injection context required
by the evaluator. Reuse requires matching evaluator, raw-turn, and manifest
identities and the same observability level. Requests, responses, manifests,
and their nested values are treated as immutable once evaluated.
"""
if not run.raw_turns:
return None

record = run.latest_online_evaluation
if (
record is not None
and record.evaluator is evaluator
and _matches_final_trace_context(context=record.context, run=run)
):
return replace(
record.result,
evidence=safe_str_list(value=record.result.evidence),
undetermined_operands=safe_str_list(
value=record.result.undetermined_operands,
),
)

context = _evaluation_context(
raw_turns=run.raw_turns,
observability_level=run.observability_level,
manifest=run.manifest,
)
return await evaluator.evaluate_async(context=context)
Loading
Loading