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
27 changes: 27 additions & 0 deletions doc/code/scoring/5_tool_call_scorer.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -394,6 +394,33 @@
"The caller owns instrumentation and trace retrieval. Automatic request\n",
"correlation does not install a remote collector or a backend adapter."
]
},
{
"cell_type": "markdown",
"id": "14",
"metadata": {},
"source": [
"## Score tool calls without traces\n",
"\n",
"`MessageToolCallScorer` reads the same `ToolsCalled` condition from stored\n",
"messages instead of spans. A tool counts only when a model-authored\n",
"`function_call` piece is paired, by call ID, with a later `function_call_output`\n",
"piece in the conversation through the scored response. `OpenAIResponseTarget`\n",
"records `ToolExecutionMetadata` on each output piece. The scorer uses this\n",
"dispatch status, not the returned payload: a tool-reported error still counts\n",
"as an invocation, while a dispatch failure does not. Older outputs without\n",
"metadata use a conservative fallback; payloads shaped like PyRIT dispatch\n",
"errors cannot prove invocation and remain undetermined.\n",
"Injected history in the `simulated_assistant` and `simulated_tool` roles does\n",
"not count either.\n",
"\n",
"Stored messages are partial evidence. Hosted tools, several response section\n",
"types, and targets that execute tools themselves leave no output pieces, so the\n",
"scorer returns true or undetermined and never false. `OpenAIResponseTarget`\n",
"records these pieces when it runs `custom_functions`. Combine it with\n",
"`OtelToolCallScorer` under `TrueFalseScoreAggregator.OR` when trace evidence is\n",
"also available."
]
}
],
"metadata": {
Expand Down
22 changes: 22 additions & 0 deletions doc/code/scoring/5_tool_call_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,3 +277,25 @@ async def local_agent_async(request: httpx.Request) -> httpx.Response:
# Other sources can use the same protocol with their own scorable types.
# The caller owns instrumentation and trace retrieval. Automatic request
# correlation does not install a remote collector or a backend adapter.

# %% [markdown]
# ## Score tool calls without traces
#
# `MessageToolCallScorer` reads the same `ToolsCalled` condition from stored
# messages instead of spans. A tool counts only when a model-authored
# `function_call` piece is paired, by call ID, with a later `function_call_output`
# piece in the conversation through the scored response. `OpenAIResponseTarget`
# records `ToolExecutionMetadata` on each output piece. The scorer uses this
# dispatch status, not the returned payload: a tool-reported error still counts
# as an invocation, while a dispatch failure does not. Older outputs without
# metadata use a conservative fallback; payloads shaped like PyRIT dispatch
# errors cannot prove invocation and remain undetermined.
# Injected history in the `simulated_assistant` and `simulated_tool` roles does
# not count either.
#
# Stored messages are partial evidence. Hosted tools, several response section
# types, and targets that execute tools themselves leave no output pieces, so the
# scorer returns true or undetermined and never false. `OpenAIResponseTarget`
# records these pieces when it runs `custom_functions`. Combine it with
# `OtelToolCallScorer` under `TrueFalseScoreAggregator.OR` when trace evidence is
# also available.
2 changes: 2 additions & 0 deletions pyrit/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,6 +240,7 @@
unregister_common_json_schema,
)
from pyrit.models.target.request_trace_context import RequestTraceContext
from pyrit.models.target.tool_execution_metadata import ToolExecutionMetadata

_LAZY_EXPORTS: dict[str, str] = {
"AllAvailableDatasetSize": "pyrit.models.scenario_dataset_size_estimate",
Expand All @@ -251,6 +252,7 @@
"ScenarioDatasetSizeEstimateKind": "pyrit.models.scenario_dataset_size_estimate",
"scenario_dataset_size_from_limit": "pyrit.models.scenario_dataset_size_estimate",
"RequestTraceContext": "pyrit.models.target.request_trace_context",
"ToolExecutionMetadata": "pyrit.models.target.tool_execution_metadata",
"AttackAnalyticsCell": "pyrit.models.analytics",
"AttackAnalyticsConverterDirection": "pyrit.models.analytics",
"AttackAnalyticsDimension": "pyrit.models.analytics",
Expand Down
36 changes: 36 additions & 0 deletions pyrit/models/target/tool_execution_metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

from __future__ import annotations

from typing import Any, ClassVar

from pydantic import BaseModel, ConfigDict, Field


class ToolExecutionMetadata(BaseModel):
"""Target-recorded invocation status, separate from a tool's returned data."""

METADATA_KEY: ClassVar[str] = "pyrit_tool_execution"
model_config = ConfigDict(frozen=True, extra="forbid")

invoked: bool = Field(strict=True)

def to_metadata(self) -> dict[str, Any]:
"""Return the metadata fragment for a function-call output piece."""
return {self.METADATA_KEY: self.model_dump(mode="json")}

@classmethod
def from_metadata(cls, *, metadata: dict[str, Any]) -> ToolExecutionMetadata | None:
"""
Read invocation status without interpreting tool-returned data.

Returns:
ToolExecutionMetadata | None: The recorded status, or None if absent.

Raises:
ValueError: If stored invocation metadata is malformed.
"""
if cls.METADATA_KEY not in metadata:
return None
return cls.model_validate(metadata[cls.METADATA_KEY])
69 changes: 40 additions & 29 deletions pyrit/prompt_target/openai/openai_response_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
MessagePiece,
PromptDataType,
PromptResponseError,
ToolExecutionMetadata,
)
from pyrit.models.messages.chat_message import FunctionCall
from pyrit.prompt_target.common.target_capabilities import TargetCapabilities
Expand Down Expand Up @@ -65,6 +66,12 @@ class _SerializedPiece:
placement: Literal["inline", "top_level"]


@dataclass(frozen=True)
class _ToolDispatchResult:
output: object
invoked: bool


class MessagePieceType(str, Enum):
"""Enumeration of different types of message pieces."""

Expand Down Expand Up @@ -661,8 +668,8 @@ async def _run_tool_call_loop_async(self, *, normalized_conversation: list[Messa
for tool_call_section in tool_call_sections:
tool_output = await self._execute_call_section_async(tool_call_section)
tool_piece = self._make_tool_piece(
tool_output,
tool_call_section["call_id"],
result=tool_output,
call_id=tool_call_section["call_id"],
reference_piece=message_piece,
)
tool_message = Message(message_pieces=[tool_piece])
Expand Down Expand Up @@ -899,17 +906,16 @@ def _find_pending_tool_calls(self, reply: Message) -> list[dict[str, Any]]:
calls.append(cast("dict[str, Any]", section))
return calls

async def _execute_call_section_async(self, tool_call_section: dict[str, Any]) -> object:
async def _execute_call_section_async(self, tool_call_section: dict[str, Any]) -> _ToolDispatchResult:
"""
Execute a function call using the matching tool.

Args:
tool_call_section: The function_call section dict.

Returns:
A JSON-serializable payload that will be sent as function_call_output.
If fail_on_missing_function=False and a function is missing or no function is not called, returns:
{"error": "function_not_found", "missing_function": "<name>", "available_functions": [...]}
The output payload and whether the tool was invoked. Tolerant dispatch failures
retain their error payload with invoked=False.

Raises:
ValueError: If the function call section is missing a 'name' field.
Expand All @@ -920,10 +926,10 @@ async def _execute_call_section_async(self, tool_call_section: dict[str, Any]) -
if not name:
if self._fail_on_missing_function:
raise ValueError("Function call section missing 'name' field")
return {
"error": "missing_function_name",
"tool_call_section": tool_call_section,
}
return _ToolDispatchResult(
output={"error": "missing_function_name", "tool_call_section": tool_call_section},
invoked=False,
)

args_json = tool_call_section.get("arguments", "{}")
try:
Expand All @@ -933,50 +939,54 @@ async def _execute_call_section_async(self, tool_call_section: dict[str, Any]) -
if self._fail_on_missing_function:
raise ValueError(f"Malformed arguments for function '{name}': {args_json}") from None
logger.warning("Malformed arguments for function '%s': %s", name, args_json)
return {
"error": "malformed_arguments",
"function": name,
"raw_arguments": args_json,
}
return _ToolDispatchResult(
output={"error": "malformed_arguments", "function": name, "raw_arguments": args_json},
invoked=False,
)
await self._initialize_tools_async()
configured_tool = next((tool for tool in self._tools if tool.name == name), None)
if configured_tool is not None:
if not isinstance(args, dict):
if self._fail_on_missing_function:
raise ValueError(f"Arguments for function '{name}' must be a JSON object")
return {
"error": "malformed_arguments",
"function": name,
"raw_arguments": args_json,
}
return await configured_tool.execute_async(arguments=args)
return _ToolDispatchResult(
output={"error": "malformed_arguments", "function": name, "raw_arguments": args_json},
invoked=False,
)
return _ToolDispatchResult(output=await configured_tool.execute_async(arguments=args), invoked=True)

custom_function = self._custom_functions.get(name)
if custom_function is not None:
return await custom_function(cast("dict[str, Any]", args))
return _ToolDispatchResult(output=await custom_function(cast("dict[str, Any]", args)), invoked=True)

if self._fail_on_missing_function:
raise KeyError(f"Function '{name}' is not registered")
available_functions = sorted({*(tool.name for tool in self._tools), *self._custom_functions})
logger.warning("Function '%s' not registered. Available: %s", name, available_functions)
return {
"error": "function_not_found",
"missing_function": name,
"available_functions": available_functions,
}
return _ToolDispatchResult(
output={
"error": "function_not_found",
"missing_function": name,
"available_functions": available_functions,
},
invoked=False,
)

def _make_tool_piece(self, output: object, call_id: str, *, reference_piece: MessagePiece) -> MessagePiece:
def _make_tool_piece(
self, *, result: _ToolDispatchResult, call_id: str, reference_piece: MessagePiece
) -> MessagePiece:
"""
Create a function_call_output MessagePiece.

Args:
output: The tool output to wrap.
result: The dispatch result to record.
call_id: The call ID for the function call.
reference_piece: A reference piece to copy conversation context from.

Returns:
A MessagePiece containing the function call output.
"""
output = result.output
output_str = output if isinstance(output, str) else json.dumps(output, separators=(",", ":"))
return MessagePiece(
role="tool",
Expand All @@ -986,4 +996,5 @@ def _make_tool_piece(self, output: object, call_id: str, *, reference_piece: Mes
),
original_value_data_type="function_call_output",
conversation_id=reference_piece.conversation_id,
prompt_metadata=ToolExecutionMetadata(invoked=result.invoked).to_metadata(),
)
2 changes: 2 additions & 0 deletions pyrit/score/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@
)
from pyrit.score.true_false.local_refusal_classifier_scorer import LocalRefusalClassifierScorer
from pyrit.score.true_false.manual_scorer import ManualScorer
from pyrit.score.true_false.message_tool_call_scorer import MessageToolCallScorer
from pyrit.score.true_false.otel_tool_call_scorer import OtelToolCallScorer
from pyrit.score.true_false.prompt_shield_scorer import PromptShieldScorer
from pyrit.score.true_false.question_answer_scorer import QuestionAnswerScorer
Expand Down Expand Up @@ -207,6 +208,7 @@
"LlamaGuardScorer": "pyrit.score.true_false.llamaguard_scorer",
"MarkdownInjectionScorer": "pyrit.score.true_false.regex.markdown_injection",
"ManualScorer": "pyrit.score.true_false.manual_scorer",
"MessageToolCallScorer": "pyrit.score.true_false.message_tool_call_scorer",
"MessageScorableResolver": "pyrit.score.message_scorable_resolver",
"MessageScorable": "pyrit.score.scorable",
"MessageScorer": "pyrit.score.message_scorer",
Expand Down
Loading
Loading