Skip to content
Open
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
68 changes: 53 additions & 15 deletions tests/test_generation_asttools.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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):
Expand All @@ -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")