From b7b13ad4be021c204ad7895ab25255c6520a093e Mon Sep 17 00:00:00 2001 From: wzj1228516103 <154067345+wzj1228516103@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:17:56 +0800 Subject: [PATCH] test: add coverage for asttools helpers --- tests/test_generation_asttools.py | 68 ++++++++++++++++++++++++------- 1 file changed, 53 insertions(+), 15 deletions(-) diff --git a/tests/test_generation_asttools.py b/tests/test_generation_asttools.py index 8e3397fc..48b209d6 100644 --- a/tests/test_generation_asttools.py +++ b/tests/test_generation_asttools.py @@ -1,9 +1,10 @@ +import ast import os +import shutil import unittest from pathlib import Path -import shutil + from agentstack.generation import asttools -import ast BASE_PATH = Path(__file__).parent @@ -96,21 +97,15 @@ def test_render_node(self): def test_insert_method(self): file_path = self.project_dir / "sample.py" - with open(file_path, "w") as f: - f.write("""class TestClass: + file_path.write_text("""class TestClass: def existing(self): pass""") - file = asttools.File(file_path) - - # Test inserting method when there's no newline at the end - new_method = """ def method1(self): - pass""" class_node = asttools.find_class(file.tree, "TestClass") method_node = asttools.find_method_in_class(class_node, "existing") - start, end = file.get_node_range(method_node) - - file.insert_method(end, new_method) + _, end = file.get_node_range(method_node) + file.insert_method(end, """ def method1(self): + pass""") self.assertEqual(file.source, """class TestClass: def existing(self): pass @@ -125,7 +120,7 @@ def method1(self): return True""" class_node = asttools.find_class(file.tree, "TestClass") method_node = asttools.find_method_in_class(class_node, "method1") - start, end = file.get_node_range(method_node) + _, end = file.get_node_range(method_node) file.insert_method(end, new_method) self.assertEqual(file.source, """class TestClass: @@ -145,8 +140,7 @@ def method2(self): print('middle')""" class_node = asttools.find_class(file.tree, "TestClass") method_node = asttools.find_method_in_class(class_node, "method1") - start, end = file.get_node_range(method_node) - + _, end = file.get_node_range(method_node) file.insert_method(end, new_method) self.assertEqual(file.source, """class TestClass: def existing(self): @@ -162,3 +156,47 @@ def method2(self): return True """) + + def test_ast_query_helpers(self): + tree = ast.parse("""from os import path +import json + +@marker +class Sample: + @decorated + def run(self): + return Worker() + + def plain(self): + value = Worker() + self.client.send(value, timeout=1) + return value +""") + self.assertEqual([node.module for node in asttools.get_all_imports(tree)], ["os"]) + self.assertIsNone(asttools.find_method(tree, "missing")) + sample = asttools.find_class(tree, "Sample") + self.assertEqual(asttools.find_method(sample, "run").name, "run") + plain_method = asttools.find_method_in_class(sample, "plain") + self.assertEqual(len(asttools.find_method_calls(plain_method, "send")), 1) + self.assertEqual(len(asttools.find_method_calls(plain_method, "Worker")), 1) + self.assertEqual(asttools.find_class_with_decorator(tree, "marker"), [sample]) + self.assertEqual(asttools.find_class_with_regex(tree, r"Sam.*"), [sample]) + self.assertEqual(asttools.find_method_in_class(sample, "plain").name, "plain") + self.assertEqual([m.name for m in asttools.find_decorated_method_in_class(sample, "decorated")], ["run"]) + instantiation_tree = ast.parse("Worker = Worker()") + self.assertIsNotNone(asttools.find_class_instantiation(instantiation_tree, "Worker")) + call = ast.parse("send(timeout=1)").body[0].value + self.assertEqual(asttools.find_kwarg_in_method_call(call, "timeout").arg, "timeout") + + def test_ast_value_and_tool_helpers(self): + self.assertEqual(asttools.create_attribute("client", "send").attr, "send") + self.assertEqual(asttools.get_node_value(ast.parse("client.send").body[0].value), "send") + self.assertEqual(asttools.get_node_value(ast.Constant(value="value")), "value") + tool = asttools.create_tool_node("search") + self.assertEqual(len(asttools.find_tool_nodes(ast.List(elts=[tool]))), 1) + self.assertEqual(asttools.find_tool_nodes(ast.List(elts=[ast.Name(id="other")])), []) + file_path = self.project_dir / "remove.py" + file_path.write_text("value = 1\nother = 2\n") + file = asttools.File(file_path) + file.remove_node(file.tree.body[0]) + self.assertEqual(file.source, "\nother = 2\n")