fix(guardrails): apply the for guards to async for in the sandbox

AsyncFor was the last async node still routed to node_contents_visit, which
only recurses into children. `for a in x` is rewritten to `for a in
_getiter_(x)` and `for a, b in x` to iterate through the unpack guard; the
async spellings got neither, so they enforced strictly less than the sync
forms they mirror.

Delegating to visit_For fixes the iterator guard, but its tuple-target rewrite
emits _iter_unpack_sequence_, a plain generator that `async for` cannot
consume ("requires an object with __aiter__ method, got generator"). So the
call is retargeted to _aiter_unpack_sequence_, an async-generator helper with
the same contract: guard the iteration, then guard each element's unpacking.

Await keeps node_contents_visit -- it has no synchronous counterpart and only
wraps an expression.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_0147HaHohHPnpFBUa4tUkvPg
This commit is contained in:
lzhan011 2026-09-02 16:24:04 -05:00
parent 9dcaca2522
commit 03f6d79c43
2 changed files with 95 additions and 4 deletions

View file

@ -16,7 +16,7 @@ restriction intact.
import ast
import operator
from collections.abc import Callable, Mapping
from collections.abc import AsyncIterable, AsyncIterator, Callable, Mapping
from types import CodeType
from typing import Final
@ -48,15 +48,19 @@ class AsyncAwareTransformer(RestrictingNodeTransformer):
inherited automatically. ``AsyncWith`` gets the same treatment for the same
reason: ``node_contents_visit`` only recurses into children, so routing it
there left ``async with x as (a, b)`` without the unpack guard that
``with x as (a, b)`` gets. ``AsyncFor``/``Await`` delegate to
``node_contents_visit`` so their children still get transformed.
``with x as (a, b)`` gets. ``AsyncFor`` likewise delegates to ``visit_For``,
which is what wraps the loop iterator in ``_getiter_``. ``Await`` has no
synchronous counterpart and only wraps an expression, so it stays on
``node_contents_visit``.
"""
def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> ast.AST:
return self.visit_FunctionDef(node)
def visit_AsyncFor(self, node: ast.AsyncFor) -> ast.AST:
return self.node_contents_visit(node)
transformed: Final = self.visit_For(node)
_use_async_iter_unpack(transformed)
return transformed
def visit_AsyncWith(self, node: ast.AsyncWith) -> ast.AST:
return self.visit_With(node)
@ -65,6 +69,39 @@ class AsyncAwareTransformer(RestrictingNodeTransformer):
return self.node_contents_visit(node)
_ITER_UNPACK_NAME: Final = "_iter_unpack_sequence_"
_ASYNC_ITER_UNPACK_NAME: Final = "_aiter_unpack_sequence_"
def _use_async_iter_unpack(node: ast.AST) -> None:
"""Point a transformed ``async for`` at the async unpack guard.
``visit_For`` rewrites ``for a, b in x`` into
``for (a, b) in _iter_unpack_sequence_(x, spec, _getiter_)``. That helper is
a plain generator, which ``async for`` cannot consume, so the async form
needs the async-generator equivalent under its own name.
"""
iter_node: Final = getattr(node, "iter", None)
if (
isinstance(iter_node, ast.Call)
and isinstance(iter_node.func, ast.Name)
and iter_node.func.id == _ITER_UNPACK_NAME
):
iter_node.func.id = _ASYNC_ITER_UNPACK_NAME
async def _aiter_unpack_sequence_(
it: object, spec: object, _getiter_: Callable[[object], AsyncIterable[object]]
) -> AsyncIterator[object]:
"""``guarded_iter_unpack_sequence`` for ``async for`` targets.
Same contract as the RestrictedPython helper — guard the iteration, then
guard each element's sequence unpacking — over an async iterator.
"""
async for ob in _getiter_(it):
yield guarded_unpack_sequence(ob, spec, _getiter_)
_INPLACE_OPS: Final[Mapping[str, Callable[[object, object], object]]] = {
"+=": operator.iadd,
"-=": operator.isub,
@ -125,6 +162,7 @@ def build_sandbox_globals() -> dict[str, object]:
# ships no default. Without it those statements compile and then raise
# NameError the first time the guardrail runs.
"_unpack_sequence_": guarded_unpack_sequence,
_ASYNC_ITER_UNPACK_NAME: _aiter_unpack_sequence_,
"_write_": full_write_guard,
"_inplacevar_": _inplacevar_,
}

View file

@ -311,3 +311,56 @@ async def test_async_with_still_executes():
sandbox_globals,
)
assert await sandbox_globals["f"](_ACtx((5, 6))) == 11
def test_async_for_gets_the_same_guards_as_for():
"""`async for` is the async spelling of `for`; it must not enforce less."""
sync_guards = _guard_names("def f(x):\n for a in x:\n pass\n")
async_guards = _guard_names("async def f(x):\n async for a in x:\n pass\n")
assert "_getiter_" in sync_guards, "precondition: `for` iteration is guarded"
assert sync_guards <= async_guards
def test_async_for_unpacking_uses_the_async_unpack_guard():
"""Tuple targets are guarded too, via the async-iterable variant of the helper."""
guards = _guard_names("async def f(x):\n async for a, b in x:\n pass\n")
assert "_aiter_unpack_sequence_" in guards
# The sync generator would raise "requires an object with __aiter__".
assert "_iter_unpack_sequence_" not in guards
@pytest.mark.asyncio
@pytest.mark.parametrize(
("body", "items", "expected"),
[
(" async for i in src:\n out.append(i)\n", [1, 2, 3], [1, 2, 3]),
(" async for a, b in src:\n out.append(a + b)\n", [(1, 2), (3, 4)], [3, 7]),
],
)
async def test_async_for_still_executes(body: str, items: list, expected: list):
"""Guarding `async for` must not break it, with or without a tuple target."""
from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
build_sandbox_globals,
compile_sandboxed,
)
class _AIter:
def __init__(self, values):
self.values = list(values)
def __aiter__(self):
return self
async def __anext__(self):
if not self.values:
raise StopAsyncIteration
return self.values.pop(0)
sandbox_globals = build_sandbox_globals()
exec( # noqa: S102
compile_sandboxed("async def f(src):\n out = []\n" + body + " return out\n"),
sandbox_globals,
)
assert await sandbox_globals["f"](_AIter(items)) == expected