diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py index 1588836f681..d5707894a9d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/sandbox.py @@ -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) diff --git a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py index d245096764d..679c0966828 100644 --- a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py +++ b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py @@ -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