diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index b7d87a5f43c..b5e080eb83c 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -73,6 +73,7 @@ "get_security_tools", "inspect_variable", "quarantined_llm", + "rewritten_arguments", "set_quarantine_client", "store_untrusted_content", ] @@ -101,6 +102,7 @@ # ``variable_ids`` list internally. Expanding their arguments would replace the ID # with the content and break the lookup. _VARIABLE_ID_CONSUMERS = frozenset({"inspect_variable", "quarantined_llm"}) +_REWRITTEN_ARGUMENT_INDICES_KEY = "_rewritten_argument_indices" def _get_additional_properties(obj: Any) -> dict[str, Any]: @@ -109,6 +111,23 @@ def _get_additional_properties(obj: Any) -> dict[str, Any]: return cast(dict[str, Any], props) if isinstance(props, dict) else {} +def _top_level_argument_value(context: FunctionInvocationContext, arg_name: str) -> tuple[Any, str | None]: + """Locate a top-level argument value in either context.arguments or context.kwargs. + + Returns the value and a string indicating its source ('arguments' or 'kwargs'), + or (None, None) if not found. + """ + args = cast(Any, context.arguments) + if isinstance(args, Mapping) and arg_name in args: + return cast(Any, args[arg_name]), "arguments" + + kwargs = cast(Any, context.kwargs) + if isinstance(kwargs, Mapping) and arg_name in kwargs: + return cast(Any, kwargs[arg_name]), "kwargs" + + return None, None + + @dataclass(frozen=True, order=True, slots=True) class _Principal: """Canonical tenant/user identity used internally for comparisons.""" @@ -1252,6 +1271,11 @@ def list_variables(self) -> list[str]: default=None, ) +_current_context: ContextVar[FunctionInvocationContext | None] = ContextVar( + "agent_framework_current_security_context", + default=None, +) + @experimental(feature_id=ExperimentalFeature.FIDES) class LabelTrackingFunctionMiddleware(FunctionMiddleware, _SecurityScopeBinding): @@ -1391,6 +1415,29 @@ def _update_context_label(self, new_content_label: ContentLabel) -> None: self._context_label.confidentiality.value, ) + @staticmethod + def _degrade_rewritten_arguments(context: FunctionInvocationContext) -> None: + """Degrade rewritten argument indices when validation mutates arguments. + + For list arguments, marks every final list position as rewritten + (fail-closed) rather than using -1, which callers may misinterpret as + "non-list argument." A validator that moves hidden content from index 0 + to index 1 must not let a tool treat index 1 as safe. + + For non-list arguments, uses -1 as before. + """ + rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) + if not rewritten: + return + degraded: dict[str, set[int]] = {} + for arg_name in cast(dict[str, set[int]], rewritten): + value, _ = _top_level_argument_value(context, arg_name) + if isinstance(value, (list, tuple)): + degraded[arg_name] = set(range(len(cast(Sequence[Any], value)))) + else: + degraded[arg_name] = {-1} + context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = degraded + @staticmethod def _extract_primary_tool_content(expanded_content: Any, *, from_quarantined_llm: bool) -> Any: """Return the primary response from a proven quarantined LLM result. @@ -1478,6 +1525,8 @@ def _resolve_string( depth: int, active_variables: set[str], reference_count: list[int], + rewritten_paths: set[tuple[str | int, ...]] | None = None, + current_path: tuple[str | int, ...] = (), ) -> Any: if not _EMBEDDED_VAR_REF_RE.search(value): return value @@ -1496,6 +1545,8 @@ def _resolve_string( return value if whole.group("bare"): logger.warning(_BARE_REFERENCE_WARNING) + if rewritten_paths is not None: + rewritten_paths.add(current_path) return resolved def replace(match: re.Match[str]) -> str: @@ -1511,6 +1562,8 @@ def replace(match: re.Match[str]) -> str: return match.group(0) if match.group("bare"): logger.warning(_BARE_REFERENCE_WARNING) + if rewritten_paths is not None: + rewritten_paths.add(current_path) return str(resolved) return _EMBEDDED_VAR_REF_RE.sub(replace, value) @@ -1523,6 +1576,8 @@ def _resolve_value( depth: int, active_variables: set[str], reference_count: list[int], + rewritten_paths: set[tuple[str | int, ...]] | None = None, + current_path: tuple[str | int, ...] = (), ) -> Any: if isinstance(value, str): return self._resolve_string( @@ -1531,6 +1586,8 @@ def _resolve_value( depth=depth, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, + current_path=current_path, ) if isinstance(value, BaseModel): value = value.model_dump() @@ -1543,6 +1600,8 @@ def _resolve_value( depth=depth, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, + current_path=(*current_path, key), ) for key, item in value_dict.items() } @@ -1554,8 +1613,10 @@ def _resolve_value( depth=depth, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, + current_path=(*current_path, index), ) - for item in cast(list[Any], value) + for index, item in enumerate(cast(list[Any], value)) ] if isinstance(value, tuple): return tuple( @@ -1565,8 +1626,10 @@ def _resolve_value( depth=depth, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, + current_path=(*current_path, index), ) - for item in cast(tuple[Any, ...], value) + for index, item in enumerate(cast(tuple[Any, ...], value)) ) return value @@ -1578,6 +1641,8 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont labels: list[ContentLabel] = [] active_variables: set[str] = set() reference_count = [0] + rewritten_paths: set[tuple[str | int, ...]] = set() + kwargs_rewritten_paths: set[tuple[str | int, ...]] = set() if context.arguments: context.arguments = self._resolve_value( context.arguments, @@ -1585,6 +1650,7 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont depth=0, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, ) if context.kwargs: context.kwargs = cast( @@ -1595,8 +1661,48 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont depth=0, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=kwargs_rewritten_paths, ), ) + + args_mapping = cast(Any, context.arguments) + arg_names_in_arguments: set[str] = ( + set(cast(dict[str, Any], args_mapping).keys()) if isinstance(args_mapping, Mapping) else set() + ) + + rewritten_args: dict[str, set[int]] = {} + for path in rewritten_paths | kwargs_rewritten_paths: + if not path or not isinstance(path[0], str): + continue + + arg_name = path[0] + if path not in rewritten_paths and arg_name in arg_names_in_arguments: + continue + + if arg_name not in rewritten_args: + rewritten_args[arg_name] = set() + + arg_value, arg_source = _top_level_argument_value(context, arg_name) + + if arg_source is None: + rewritten_args[arg_name].add(-1) + elif len(path) == 1: + if isinstance(arg_value, (list, tuple)): + resolved_list = cast(Sequence[Any], arg_value) + rewritten_args[arg_name].update(range(len(resolved_list))) + else: + rewritten_args[arg_name].add(-1) + elif ( + len(path) > 1 + and isinstance(path[1], int) + and not isinstance(path[1], bool) + and isinstance(arg_value, (list, tuple)) + ): + rewritten_args[arg_name].add(path[1]) + else: + rewritten_args[arg_name].add(-1) + + context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = rewritten_args return labels def _get_input_labels(self, context: FunctionInvocationContext) -> list[ContentLabel]: @@ -1756,6 +1862,7 @@ async def process( """Resolve hidden arguments, publish their labels, and label the result.""" scope_token = self._activate_security_scope(context) middleware_token = _current_middleware.set(self) + context_token = _current_context.set(context) try: function_name = context.function.name if "original_arguments_for_messages" not in context.metadata: @@ -1778,14 +1885,16 @@ async def process( boundary="security policy", ) if context.metadata.get(_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY) is True: + pre_validation_token = _argument_authority_token(context.arguments, boundary="security policy") context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] context, context.arguments, ) - context.metadata[_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY] = _argument_authority_token( - context.arguments, - boundary="security policy", - ) + post_validation_token = _argument_authority_token(context.arguments, boundary="security policy") + context.metadata[_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY] = post_validation_token + if pre_validation_token != post_validation_token: + self._degrade_rewritten_arguments(context) + argument_labels = [*input_labels, *resolved_labels] argument_label = combine_labels(*argument_labels) if argument_labels else ContentLabel() @@ -1839,6 +1948,7 @@ async def process( return self._label_result(context, function_name, fallback_label) finally: + _current_context.reset(context_token) _current_middleware.reset(middleware_token) self._active_security_scope.reset(scope_token) @@ -2245,6 +2355,51 @@ def get_current_middleware() -> LabelTrackingFunctionMiddleware | None: return _current_middleware.get() +def rewritten_arguments(context: FunctionInvocationContext | None = None) -> dict[str, set[int]]: + """Get a mapping of argument names to the set of rewritten positions. + + Returns a dictionary where keys are argument names and values are sets of + indices. For list arguments, the set contains the indices of the items + that were rewritten by variable expansion. For non-list arguments, the set + contains -1. + + If the arguments were mutated after the indices were published (e.g., by a + Pydantic validator that reorders or filters a list), all final list + positions are marked as rewritten (fail-closed) because the original + positions are stale and a validator may have moved hidden content to any + position. For non-list arguments, -1 is used. This covers both the + auto-preparation path and the public direct-middleware path. + + Args: + context: The function invocation context. If None, the context from + the current execution flow is used. + + Returns: + A dictionary mapping argument names to sets of rewritten indices. + """ + if context is None: + context = _current_context.get() + if context is None: + return {} + rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) + if rewritten is None: + return {} + + indices_snapshot = context.metadata.get(_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY) + if indices_snapshot is not None: + current_snapshot = _argument_authority_token(context.arguments, boundary="security policy") + if current_snapshot != indices_snapshot: + result: dict[str, set[int]] = {} + for arg_name in cast(dict[str, set[int]], rewritten): + value, _ = _top_level_argument_value(context, arg_name) + if isinstance(value, (list, tuple)): + result[arg_name] = set(range(len(cast(Sequence[Any], value)))) + else: + result[arg_name] = {-1} + return result + return {k: set(v) for k, v in cast(dict[str, set[int]], rewritten).items()} + + @dataclass(frozen=True, slots=True) class _PendingPolicyApproval: """Immutable binding record for a pending policy-violation approval. diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index 7f371b61a1d..9d341d70572 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -8,11 +8,11 @@ import math from datetime import timedelta from types import MappingProxyType, SimpleNamespace -from typing import Any, cast +from typing import Annotated, Any, cast from unittest.mock import AsyncMock import pytest -from pydantic import BaseModel, field_validator +from pydantic import AfterValidator, BaseModel, field_validator from agent_framework import ( Agent, @@ -26,6 +26,7 @@ ) from agent_framework._middleware import FunctionMiddlewarePipeline, MiddlewareFailure, MiddlewareTermination from agent_framework._tools import ( + _AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY, # pyright: ignore[reportPrivateUsage] FunctionTool, _auto_invoke_function, _resolve_approval_responses, @@ -46,6 +47,7 @@ VariableReferenceContent, combine_labels, get_current_middleware, + rewritten_arguments, store_untrusted_content, ) @@ -7569,3 +7571,506 @@ async def execute(_context: FunctionInvocationContext) -> list[Content]: assert executed is True assert replay.metadata["user_approved_violation"] is True + + +@pytest.mark.asyncio +async def test_rewritten_arguments_no_rewrites(): + """Verify normal/non-expanded path returns empty dict.""" + tracker = LabelTrackingFunctionMiddleware() + captured: dict[str, Any] = {} + + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool( + name="my_tool", + func=my_tool, + additional_properties={"accepts_untrusted": True}, + ) + + context = FunctionInvocationContext( + function=tool, + arguments={"files": ["one.txt", "two.txt"]}, + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_explicit_context(): + """Verify the explicit context API works inside a tool.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + captured_context: FunctionInvocationContext | None = None + + async def my_tool(files: list[str]): + captured["explicit"] = rewritten_arguments(captured_context) + captured["implicit"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext(function=tool, arguments={"files": [f"[{var_id}]", "safe.txt"]}) + captured_context = context + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["explicit"] == {"files": {0}} + assert captured["implicit"] == {"files": {0}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_multiple_args(): + """Verify tracking works across multiple top-level arguments.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id1 = store.store("file_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + var_id2 = store.store("msg_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(files: list[str], message: str): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, arguments={"files": [f"[{var_id1}]", "safe.txt"], "message": f"[{var_id2}]"} + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {"files": {0}, "message": {-1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_duplicate_equal_values(): + """Test that duplicate/equal final values are tracked per position.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + + var_id1 = store.store("same_string", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + var_id2 = store.store("same_string", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, arguments={"files": [f"[{var_id1}]", f"[{var_id2}]", "normal.txt"]} + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["received"] == ["same_string", "same_string", "normal.txt"] + assert captured["rewritten"] == {"files": {0, 1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_multiple_list_positions(): + """Test multiple list positions alongside untouched positions.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, arguments={"files": [f"[{var_id}]", "untouched.txt", f"[{var_id}]"]} + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {"files": {0, 2}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_scalar(): + """Test scalar (non-list) arguments.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(text: str): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext(function=tool, arguments={"text": f"[{var_id}]"}) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {"text": {-1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_nested_dict_semantics(): + """Test that nested rewrites are reported against the top-level argument.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(config: dict): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext(function=tool, arguments={"config": {"path": f"[{var_id}]", "safe": "txt"}}) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {"config": {-1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_whole_list_substitution(): + """Test that whole-list substitution reports all resulting indices.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store(["hidden1.txt", "hidden2.txt"], ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext(function=tool, arguments={"files": f"[{var_id}]"}) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["received"] == ["hidden1.txt", "hidden2.txt"] + assert captured["rewritten"] == {"files": {0, 1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_integer_keyed_dict(): + """Test that integer-keyed dictionaries are not treated as lists.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(config: dict): + captured["rewritten"] = rewritten_arguments() + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"config": {0: f"[{var_id}]", 1: "safe.txt"}}, + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["rewritten"] == {"config": {-1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_asyncio_to_thread(): + """Verify async/thread/context behavior with asyncio.to_thread.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("thread_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def threaded_tool(files: list[str]): + def worker(): + return rewritten_arguments() + + captured["rewritten"] = await asyncio.to_thread(worker) + return "ok" + + tool = FunctionTool(name="threaded_tool", func=threaded_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext(function=tool, arguments={"files": [f"[{var_id}]", "safe.txt"]}) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + + assert captured["rewritten"] == {"files": {0}} + + +def _reorder_files(v: list[str]) -> list[str]: + """A validator that reorders the list, invalidating original indices.""" + return sorted(v, reverse=True) + + +def _preserve_files(v: list[str]) -> list[str]: + """A validator that runs but does not mutate the list.""" + return v + + +ReorderedFiles = Annotated[list[str], AfterValidator(_reorder_files)] +PreservedFiles = Annotated[list[str], AfterValidator(_preserve_files)] + + +@pytest.mark.asyncio +async def test_rewritten_arguments_degrade_when_validator_reorders(): + """Verify indices degrade to -1 when a Pydantic validator mutates a list.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + + var_id = store.store("a", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + + captured: dict[str, Any] = {} + + async def my_tool(files: ReorderedFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "c", "b"]}, + metadata={_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY: True}, + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["c", "b", "a"] + assert captured["rewritten"] == {"files": {0, 1, 2}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_preserved_when_validator_does_not_mutate(): + """Verify indices are preserved when a validator runs but doesn't mutate the list.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: PreservedFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "safe.txt"]}, + metadata={_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY: True}, + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["payload", "safe.txt"] + assert captured["rewritten"] == {"files": {0}} + + +def _filter_files(v: list[str]) -> list[str]: + """A validator that filters out items, invalidating original indices.""" + return [f for f in v if f != "drop_me"] + + +FilteredFiles = Annotated[list[str], AfterValidator(_filter_files)] + + +@pytest.mark.asyncio +async def test_rewritten_arguments_degrade_when_validator_filters(): + """Verify indices degrade to -1 when a Pydantic validator filters a list.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("keep_me", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: FilteredFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "drop_me", "safe.txt"]}, + metadata={_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY: True}, + ) + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["keep_me", "safe.txt"] + assert captured["rewritten"] == {"files": {0, 1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_degrade_when_validator_reorders_without_auto_prep(): + """Indices degrade to -1 when a Pydantic validator reorders a list, + even without _AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY (public direct-middleware path).""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("a", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: ReorderedFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "c", "b"]}, + ) + + async def call_next(): + context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] + context, + context.arguments, + ) + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["c", "b", "a"] + assert captured["rewritten"] == {"files": {0, 1, 2}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_degrade_when_validator_filters_without_auto_prep(): + """Indices degrade to -1 when a Pydantic validator filters a list, + even without _AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("keep_me", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: FilteredFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "drop_me", "safe.txt"]}, + ) + + async def call_next(): + context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] + context, + context.arguments, + ) + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["keep_me", "safe.txt"] + assert captured["rewritten"] == {"files": {0, 1}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_preserved_when_validator_does_not_mutate_without_auto_prep(): + """Indices are preserved when a validator runs but doesn't mutate, + even without _AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: PreservedFiles): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": [f"[{var_id}]", "safe.txt"]}, + ) + + async def call_next(): + context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] + context, + context.arguments, + ) + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["payload", "safe.txt"] + assert captured["rewritten"] == {"files": {0}} + + +@pytest.mark.asyncio +async def test_rewritten_arguments_kwargs_collision_excluded(): + """Rewrites in context.kwargs for a name also in context.arguments + must not be published — _top_level_argument_value prefers + context.arguments, so a kwargs rewrite for a shared name would + report indices against the untouched callable argument.""" + tracker = LabelTrackingFunctionMiddleware() + store = tracker.get_variable_store() + var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + captured: dict[str, Any] = {} + + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files + return "ok" + + tool = FunctionTool(name="my_tool", func=my_tool, additional_properties={"accepts_untrusted": True}) + context = FunctionInvocationContext( + function=tool, + arguments={"files": ["safe.txt", "also_safe.txt"]}, + ) + context.kwargs = {"files": [f"[{var_id}]", "safe.txt"]} + + async def call_next(): + await tool.invoke(arguments=context.arguments) + + await tracker.process(context, call_next) + assert captured["received"] == ["safe.txt", "also_safe.txt"] + assert captured["rewritten"] == {}