From 8c58607fea25f8bc66fd6ffedcbec486d732b470 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Fri, 18 Sep 2026 15:10:02 +0530 Subject: [PATCH 1/6] feat(security): publish rewritten argument positions from variable expansion --- .../packages/core/agent_framework/security.py | 71 +++++- python/packages/core/tests/test_security.py | 214 ++++++++++++++++++ 2 files changed, 283 insertions(+), 2 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index b7d87a5f43c..496bc6fac8d 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]: @@ -1252,6 +1254,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): @@ -1478,6 +1485,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 +1505,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 +1522,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 +1536,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 +1546,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 +1560,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 +1573,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 +1586,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 +1601,7 @@ 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() if context.arguments: context.arguments = self._resolve_value( context.arguments, @@ -1585,6 +1609,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 +1620,23 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont depth=0, active_variables=active_variables, reference_count=reference_count, + rewritten_paths=rewritten_paths, ), ) + + rewritten_args: dict[str, set[int]] = {} + for path in rewritten_paths: + if not path or not isinstance(path[0], str): + continue + arg_name = path[0] + if arg_name not in rewritten_args: + rewritten_args[arg_name] = set() + if len(path) > 1 and isinstance(path[1], int): + 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 +1796,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: @@ -1839,6 +1880,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 +2287,31 @@ 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. + + 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 {} + 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..780ef6659b1 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -46,6 +46,7 @@ VariableReferenceContent, combine_labels, get_current_middleware, + rewritten_arguments, store_untrusted_content, ) @@ -7569,3 +7570,216 @@ 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 = {} + + 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 = {} + + async def my_tool(files: list[str], context: FunctionInvocationContext): + captured["explicit"] = rewritten_arguments(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"]}) + + async def call_next(): + await tool.func(files=context.arguments["files"], context=context) + + 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 = {} + + 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 = {} + + 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 = {} + + 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 = {} + + 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 = {} + + 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_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 = {} + + 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}} From 14f2c7df4c313405329e724b05b153cbdc4e03f0 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Sat, 19 Sep 2026 14:28:35 +0530 Subject: [PATCH 2/6] fix(security): treat integer-keyed dicts as non-list arguments --- .../packages/core/agent_framework/security.py | 29 ++++++++++++++++++- python/packages/core/tests/test_security.py | 26 +++++++++++++++++ 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 496bc6fac8d..526cce206be 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -111,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.""" @@ -1628,10 +1645,20 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont for path in rewritten_paths: if not path or not isinstance(path[0], str): continue + arg_name = path[0] if arg_name not in rewritten_args: rewritten_args[arg_name] = set() - if len(path) > 1 and isinstance(path[1], int): + + arg_value, arg_source = _top_level_argument_value(context, arg_name) + + if ( + arg_source is not None + and 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) diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index 780ef6659b1..4d41afa1fd4 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -7783,3 +7783,29 @@ async def call_next(): await tracker.process(context, call_next) assert captured["rewritten"] == {"files": {0}} + + +@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 = {} + + 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}} From b6859deeb435e4c996040d5b14e67ed32a22cfd1 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Mon, 21 Sep 2026 23:46:24 +0530 Subject: [PATCH 3/6] fix(security): handle whole-list substitution and fix typing --- .../packages/core/agent_framework/security.py | 13 +++- python/packages/core/tests/test_security.py | 74 +++++++++++++------ 2 files changed, 61 insertions(+), 26 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 526cce206be..72ea03fc079 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -1652,9 +1652,16 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont arg_value, arg_source = _top_level_argument_value(context, arg_name) - if ( - arg_source is not None - and len(path) > 1 + 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)) diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index 4d41afa1fd4..db9909e27b1 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -7576,7 +7576,7 @@ async def execute(_context: FunctionInvocationContext) -> list[Content]: async def test_rewritten_arguments_no_rewrites(): """Verify normal/non-expanded path returns empty dict.""" tracker = LabelTrackingFunctionMiddleware() - captured = {} + captured: dict[str, Any] = {} async def my_tool(files: list[str]): captured["rewritten"] = rewritten_arguments() @@ -7608,18 +7608,20 @@ async def test_rewritten_arguments_explicit_context(): store = tracker.get_variable_store() var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} + captured_context: FunctionInvocationContext | None = None - async def my_tool(files: list[str], context: FunctionInvocationContext): - captured["explicit"] = rewritten_arguments(context) + 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.func(files=context.arguments["files"], context=context) + await tool.invoke(arguments=context.arguments) await tracker.process(context, call_next) @@ -7635,7 +7637,7 @@ async def test_rewritten_arguments_multiple_args(): var_id1 = store.store("file_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) var_id2 = store.store("msg_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(files: list[str], message: str): captured["rewritten"] = rewritten_arguments() @@ -7663,7 +7665,7 @@ async def test_rewritten_arguments_duplicate_equal_values(): var_id1 = store.store("same_string", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) var_id2 = store.store("same_string", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(files: list[str]): captured["rewritten"] = rewritten_arguments() @@ -7691,7 +7693,7 @@ async def test_rewritten_arguments_multiple_list_positions(): store = tracker.get_variable_store() var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(files: list[str]): captured["rewritten"] = rewritten_arguments() @@ -7717,7 +7719,7 @@ async def test_rewritten_arguments_scalar(): store = tracker.get_variable_store() var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(text: str): captured["rewritten"] = rewritten_arguments() @@ -7741,7 +7743,7 @@ async def test_rewritten_arguments_nested_dict_semantics(): store = tracker.get_variable_store() var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(config: dict): captured["rewritten"] = rewritten_arguments() @@ -7759,30 +7761,29 @@ async def call_next(): @pytest.mark.asyncio -async def test_rewritten_arguments_asyncio_to_thread(): - """Verify async/thread/context behavior with asyncio.to_thread.""" +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("thread_content", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) + var_id = store.store(["hidden1.txt", "hidden2.txt"], ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} - async def threaded_tool(files: list[str]): - def worker(): - return rewritten_arguments() - - captured["rewritten"] = await asyncio.to_thread(worker) + async def my_tool(files: list[str]): + captured["rewritten"] = rewritten_arguments() + captured["received"] = files 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"]}) + 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["rewritten"] == {"files": {0}} + assert captured["received"] == ["hidden1.txt", "hidden2.txt"] + assert captured["rewritten"] == {"files": {0, 1}} @pytest.mark.asyncio @@ -7792,7 +7793,7 @@ async def test_rewritten_arguments_integer_keyed_dict(): store = tracker.get_variable_store() var_id = store.store("payload", ContentLabel(integrity=IntegrityLabel.UNTRUSTED)) - captured = {} + captured: dict[str, Any] = {} async def my_tool(config: dict): captured["rewritten"] = rewritten_arguments() @@ -7809,3 +7810,30 @@ async def call_next(): 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}} From 90604ec0051bfc983c7ce027ac4d70f63487971f Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Thu, 24 Sep 2026 00:02:00 +0530 Subject: [PATCH 4/6] fix(security): Refresh rewritten indices after Pydantic validation --- .../packages/core/agent_framework/security.py | 38 ++++++ python/packages/core/tests/test_security.py | 114 +++++++++++++++++- 2 files changed, 150 insertions(+), 2 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 72ea03fc079..4fd00951642 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -128,6 +128,42 @@ def _top_level_argument_value(context: FunctionInvocationContext, arg_name: str) return None, None +def _snapshot_arguments(arguments: object) -> dict[str, Any]: + """Return a deep copy of mapping arguments, or an empty dict otherwise.""" + if isinstance(arguments, Mapping): + return deepcopy(dict(cast("Mapping[str, Any]", arguments))) + return {} + + +def _refresh_rewritten_indices( + context: FunctionInvocationContext, + pre_validation_snapshot: Mapping[str, Any], +) -> None: + """Refresh rewritten argument indices after Pydantic validation. + + This function detects mutations and degrades the affected arguments' indices + to the scalar sentinel ({-1}) if the value changed. + """ + rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) + if not rewritten: + return + + refreshed: dict[str, set[int]] = {} + for arg_name, indices in cast(dict[str, set[int]], rewritten).items(): + pre_val = pre_validation_snapshot.get(arg_name) + cur_val, _ = _top_level_argument_value(context, arg_name) + + if cur_val is None: + continue + + if cur_val != pre_val: + refreshed[arg_name] = {-1} + else: + refreshed[arg_name] = set(indices) + + context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = refreshed + + @dataclass(frozen=True, order=True, slots=True) class _Principal: """Canonical tenant/user identity used internally for comparisons.""" @@ -1853,6 +1889,7 @@ async def process( boundary="security policy", ) if context.metadata.get(_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY) is True: + pre_validation_args = _snapshot_arguments(context.arguments) context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] context, context.arguments, @@ -1861,6 +1898,7 @@ async def process( context.arguments, boundary="security policy", ) + _refresh_rewritten_indices(context, pre_validation_args) argument_labels = [*input_labels, *resolved_labels] argument_label = combine_labels(*argument_labels) if argument_labels else ContentLabel() diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index db9909e27b1..f2bbfccda9d 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, @@ -7837,3 +7838,112 @@ async def call_next(): 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": {-1}} + + +@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": {-1}} From 7f92e31cfbc23abfd494b346c9595ccdd7659caa Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Fri, 25 Sep 2026 12:34:40 +0530 Subject: [PATCH 5/6] fix(security): Use type-aware token for rewritten indices comparison --- .../packages/core/agent_framework/security.py | 46 ++++--------------- 1 file changed, 8 insertions(+), 38 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 4fd00951642..0fb725910c0 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -128,42 +128,6 @@ def _top_level_argument_value(context: FunctionInvocationContext, arg_name: str) return None, None -def _snapshot_arguments(arguments: object) -> dict[str, Any]: - """Return a deep copy of mapping arguments, or an empty dict otherwise.""" - if isinstance(arguments, Mapping): - return deepcopy(dict(cast("Mapping[str, Any]", arguments))) - return {} - - -def _refresh_rewritten_indices( - context: FunctionInvocationContext, - pre_validation_snapshot: Mapping[str, Any], -) -> None: - """Refresh rewritten argument indices after Pydantic validation. - - This function detects mutations and degrades the affected arguments' indices - to the scalar sentinel ({-1}) if the value changed. - """ - rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) - if not rewritten: - return - - refreshed: dict[str, set[int]] = {} - for arg_name, indices in cast(dict[str, set[int]], rewritten).items(): - pre_val = pre_validation_snapshot.get(arg_name) - cur_val, _ = _top_level_argument_value(context, arg_name) - - if cur_val is None: - continue - - if cur_val != pre_val: - refreshed[arg_name] = {-1} - else: - refreshed[arg_name] = set(indices) - - context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = refreshed - - @dataclass(frozen=True, order=True, slots=True) class _Principal: """Canonical tenant/user identity used internally for comparisons.""" @@ -1889,7 +1853,7 @@ async def process( boundary="security policy", ) if context.metadata.get(_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY) is True: - pre_validation_args = _snapshot_arguments(context.arguments) + pre_validation_token = _argument_authority_token(context.arguments, boundary="security policy") context.function._prepare_context_arguments( # pyright: ignore[reportPrivateUsage] context, context.arguments, @@ -1898,7 +1862,13 @@ async def process( context.arguments, boundary="security policy", ) - _refresh_rewritten_indices(context, pre_validation_args) + if pre_validation_token != context.metadata[_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY]: + rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) + if rewritten: + context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = { + arg_name: {-1} for arg_name in cast(dict[str, set[int]], rewritten) + } + argument_labels = [*input_labels, *resolved_labels] argument_label = combine_labels(*argument_labels) if argument_labels else ContentLabel() From 22edce28a62243cb7bfe9ae030c1e780d135865c Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Tue, 29 Sep 2026 12:08:09 +0530 Subject: [PATCH 6/6] harden rewritten-argument indices against validator mutation and kwargs collision --- .../packages/core/agent_framework/security.py | 70 ++++++++-- python/packages/core/tests/test_security.py | 131 +++++++++++++++++- 2 files changed, 187 insertions(+), 14 deletions(-) diff --git a/python/packages/core/agent_framework/security.py b/python/packages/core/agent_framework/security.py index 0fb725910c0..b5e080eb83c 100644 --- a/python/packages/core/agent_framework/security.py +++ b/python/packages/core/agent_framework/security.py @@ -1415,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. @@ -1619,6 +1642,7 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont 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, @@ -1637,16 +1661,24 @@ def _expand_variable_references_in_context(self, context: FunctionInvocationCont depth=0, active_variables=active_variables, reference_count=reference_count, - rewritten_paths=rewritten_paths, + 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: + 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() @@ -1858,16 +1890,10 @@ async def process( context, context.arguments, ) - context.metadata[_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY] = _argument_authority_token( - context.arguments, - boundary="security policy", - ) - if pre_validation_token != context.metadata[_SECURITY_ARGUMENTS_SNAPSHOT_CONTEXT_KEY]: - rewritten = context.metadata.get(_REWRITTEN_ARGUMENT_INDICES_KEY) - if rewritten: - context.metadata[_REWRITTEN_ARGUMENT_INDICES_KEY] = { - arg_name: {-1} for arg_name in cast(dict[str, set[int]], rewritten) - } + 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() @@ -2337,6 +2363,13 @@ def rewritten_arguments(context: FunctionInvocationContext | None = None) -> dic 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. @@ -2351,6 +2384,19 @@ def rewritten_arguments(context: FunctionInvocationContext | None = None) -> dic 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()} diff --git a/python/packages/core/tests/test_security.py b/python/packages/core/tests/test_security.py index f2bbfccda9d..9d341d70572 100644 --- a/python/packages/core/tests/test_security.py +++ b/python/packages/core/tests/test_security.py @@ -7881,7 +7881,7 @@ async def call_next(): await tracker.process(context, call_next) assert captured["received"] == ["c", "b", "a"] - assert captured["rewritten"] == {"files": {-1}} + assert captured["rewritten"] == {"files": {0, 1, 2}} @pytest.mark.asyncio @@ -7946,4 +7946,131 @@ async def call_next(): await tracker.process(context, call_next) assert captured["received"] == ["keep_me", "safe.txt"] - assert captured["rewritten"] == {"files": {-1}} + 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"] == {}