Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions src/smolagents/tool_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,11 +57,14 @@ def visit_Assign(self, node):
self.visit(node.value)

def visit_With(self, node):
"""Track aliases in 'with' statements (the 'y' in 'with X as y')"""
"""Track names bound by 'with' aliases, including unpacking targets."""
for item in node.items:
if item.optional_vars: # This is the 'y' in 'with X as y'
if isinstance(item.optional_vars, ast.Name):
self.assigned_names.add(item.optional_vars.id)
self.assigned_names.update(
target.id
for target in ast.walk(item.optional_vars)
if isinstance(target, ast.Name) and isinstance(target.ctx, ast.Store)
)
self.generic_visit(node)

def visit_ExceptHandler(self, node):
Expand Down
79 changes: 79 additions & 0 deletions tests/test_tool_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,39 @@ def test_tool_to_dict_validation_with_multiple_assignments(tool_instance):
tool_instance.to_dict()


class WithUnpackingTool(Tool):
name = "with_unpacking_tool"
description = "Sum a context-managed pair"
inputs = {"value": {"type": "integer", "description": "First pair element"}}
output_type = "integer"

def forward(self, value: int) -> int:
from contextlib import nullcontext

with nullcontext((value, 3)) as (left, right):
return left + right


@tool
def tool_function_with_unpacking(value: int) -> int:
"""Sum a context-managed pair.

Args:
value: First pair element.
"""
from contextlib import nullcontext

with nullcontext((value, 3)) as (left, right):
return left + right


@pytest.mark.parametrize("tool_instance", [WithUnpackingTool(), tool_function_with_unpacking])
def test_tool_with_unpacking_roundtrip(tool_instance):
assert tool_instance(4) == 7
restored_tool = Tool.from_dict(tool_instance.to_dict())
assert restored_tool(4) == 7


class TestMethodChecker:
def test_multiple_assignments(self):
source_code = dedent(
Expand All @@ -187,3 +220,49 @@ def forward(self) -> str:
method_checker = MethodChecker(set())
method_checker.visit(ast.parse(source_code))
assert method_checker.errors == []

@pytest.mark.parametrize(
"context, result",
[
("nullcontext(7) as total", "total"),
("nullcontext((4, 3)) as (left, right)", "left + right"),
("nullcontext((4, 3)) as [left, right]", "left + right"),
("nullcontext((4, (1, 2))) as (left, [middle, right])", "left + middle + right"),
("nullcontext((4, 1, 2)) as (left, *rest)", "left + sum(rest)"),
("nullcontext((4, 1, 2)) as [left, *rest]", "left + sum(rest)"),
("nullcontext(4) as left, nullcontext(left + 3) as right", "right"),
],
)
def test_with_aliases(self, context, result):
source_code = dedent(
f"""
def forward(self):
from contextlib import nullcontext
with {context}:
return {result}
"""
)
method_checker = MethodChecker(set())
method_checker.visit(ast.parse(source_code))
assert method_checker.errors == []

@pytest.mark.parametrize(
"alias, expected_names",
[
("(left, right)", {"missing"}),
("(holder.value, right)", {"holder", "missing"}),
("(values[index], right)", {"values", "index", "missing"}),
],
)
def test_with_aliases_preserve_undefined_names(self, alias, expected_names):
source_code = dedent(
f"""
def forward(self):
from contextlib import nullcontext
with nullcontext((4, 3)) as {alias}:
return missing
"""
)
method_checker = MethodChecker(set())
method_checker.visit(ast.parse(source_code))
assert set(method_checker.errors) == {f"Name '{name}' is undefined." for name in expected_names}