mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
fix(guardrails): apply the with guards to async with in the sandbox
AsyncAwareTransformer re-permits the async nodes the RestrictedPython policy rejects outright. AsyncFunctionDef delegates to visit_FunctionDef so it inherits that visitor's checks, but AsyncWith went to node_contents_visit, which only recurses into children and performs no rewriting. The result was that the async spelling enforced less than the sync one: `with x as (a, b)` is rewritten to guard the unpacking, `async with x as (a, b)` was not. AsyncWith has the same _fields and the same security semantics as With, so delegate to visit_With the way AsyncFunctionDef delegates to visit_FunctionDef. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0147HaHohHPnpFBUa4tUkvPg
This commit is contained in:
parent
6d0146a3f3
commit
9dcaca2522
2 changed files with 59 additions and 2 deletions
|
|
@ -45,7 +45,10 @@ class AsyncAwareTransformer(RestrictingNodeTransformer):
|
|||
has the same ``_fields`` as ``FunctionDef`` and the same security
|
||||
semantics, so we delegate to ``visit_FunctionDef`` — name check, argument
|
||||
check, print-scope wrapping, and any future additions to that method are
|
||||
inherited automatically. ``AsyncFor``/``AsyncWith``/``Await`` delegate to
|
||||
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.
|
||||
"""
|
||||
|
||||
|
|
@ -56,7 +59,7 @@ class AsyncAwareTransformer(RestrictingNodeTransformer):
|
|||
return self.node_contents_visit(node)
|
||||
|
||||
def visit_AsyncWith(self, node: ast.AsyncWith) -> ast.AST:
|
||||
return self.node_contents_visit(node)
|
||||
return self.visit_With(node)
|
||||
|
||||
def visit_Await(self, node: ast.Await) -> ast.AST:
|
||||
return self.node_contents_visit(node)
|
||||
|
|
|
|||
|
|
@ -257,3 +257,57 @@ def test_tuple_unpacking_runs(body: str, expected):
|
|||
assert fn is not None
|
||||
inputs = {"pair": (1, 2), "nested": (1, (2, 3)), "seq": [1, 2, 3], "ctx": _Ctx()}
|
||||
assert fn(inputs, {}, "request") == expected
|
||||
|
||||
|
||||
def _guard_names(source: str) -> set[str]:
|
||||
"""Names of the RestrictedPython guards the compiled bytecode calls."""
|
||||
import types
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
|
||||
compile_sandboxed,
|
||||
)
|
||||
|
||||
found: set[str] = set()
|
||||
stack = [compile_sandboxed(source)]
|
||||
while stack:
|
||||
code = stack.pop()
|
||||
found |= {n for n in code.co_names + code.co_varnames if n.startswith("_") and n.endswith("_")}
|
||||
stack += [c for c in code.co_consts if isinstance(c, types.CodeType)]
|
||||
return found
|
||||
|
||||
|
||||
def test_async_with_gets_the_same_guards_as_with():
|
||||
"""`async with` is the async spelling of `with`; it must not enforce less."""
|
||||
sync_guards = _guard_names("def f(x):\n with x as (a, b):\n pass\n")
|
||||
async_guards = _guard_names("async def f(x):\n async with x as (a, b):\n pass\n")
|
||||
|
||||
assert "_unpack_sequence_" in sync_guards, "precondition: `with` unpacking is guarded"
|
||||
assert sync_guards <= async_guards
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_with_still_executes():
|
||||
"""Guarding `async with` must not break it."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
|
||||
build_sandbox_globals,
|
||||
compile_sandboxed,
|
||||
)
|
||||
|
||||
class _ACtx:
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
|
||||
async def __aenter__(self):
|
||||
return self.value
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
return False
|
||||
|
||||
sandbox_globals = build_sandbox_globals()
|
||||
exec( # noqa: S102
|
||||
compile_sandboxed(
|
||||
"async def f(ctx):\n async with ctx as (a, b):\n return a + b\n"
|
||||
),
|
||||
sandbox_globals,
|
||||
)
|
||||
assert await sandbox_globals["f"](_ACtx((5, 6))) == 11
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue