Skip to content
Merged
7 changes: 2 additions & 5 deletions agent/llm/test_generator.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,9 @@
from agent.llm.groq_client import client, model

def generate_tests(diff: str, context: list = None, intent: dict = None) -> str:

def generate_tests(diff:str,context:list =None, intent:dict = None)->str:
context_text = ""
if context:
context_text = "\n\n".join(context[:5])
context_text = "\n\n".join(context[:5]) # Include only the first 5 chunks for context

# Preparing intent text
intent_text = ""
if intent:
properties = ", ".join(intent.get("properties", []))
Expand Down
46 changes: 38 additions & 8 deletions agent/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from agent.github.committer import commit_tests
import subprocess , os
from agent.llm.groq_client import generate_review
from agent.llm.groq_client import infer_intent
from agent.github.commenter import post_comment
def get_pr_diff():
try:
Expand Down Expand Up @@ -65,31 +66,60 @@ def main():
print("Retrieving relevant context...")
relevant_chunks = query_embeddings(query_embedding)

# 7. Generate review
# 7. Infer intent
print("Inferring intent...")
intent = infer_intent(diff)

if not intent or "error" in intent:
print("Intent extraction failed, using fallback")
intent = {"purpose": "", "properties": [], "edge_cases": []}

print("Intent extracted:")
print(intent)

# 8. Generate review
print("Generating review...")
review = generate_review(diff, context=relevant_chunks)
print("Review generated:")
print(review)

# 8. Generate tests
# 9. Generate tests
print("Generating tests...")
tests = generate_tests(diff, context=relevant_chunks)
tests = generate_tests(diff, context=relevant_chunks, intent=intent)
print("Generated tests with intent:")
print(tests)
os.makedirs("tests", exist_ok=True)
with open("tests/test_generated.py", "w", encoding="utf-8") as f:
f.write(tests) # Save generated tests to a file for potential commit
print("Generated tests saved to tests/test_generated.py")

# 9. Combine output
final_output = f"{review}\n\n---\n\n### Suggested Tests\n{tests}"
# 10. Combine output
formatted_tests = f"""```bash
pytest tests/test_generated.py
{tests}
```"""

final_output = f"""{review}

---

### Suggested Tests
{formatted_tests}
"""

print("\n FINAL OUTPUT:\n")
print(final_output)

# 10. Post comment
# 11. Post comment
print("Posting comment...")
post_comment(final_output)

# 11. Commit tests
# 12. Commit tests
print("Committing tests...")
commit_tests()
try:
commit_tests()
except Exception as e:
print(f"Commit failed: {e}")

print("Done")

Expand Down
47 changes: 17 additions & 30 deletions tests/test_generated.py
Original file line number Diff line number Diff line change
@@ -1,34 +1,21 @@
import pytest
from your_module import store_embeddings, query_embeddings, client
def test_empty_diff():
assert generate_tests(diff="") == ""

def test_store_embeddings():
chunks = ["chunk1", "chunk2"]
embeddings = [[1, 2], [3, 4]]
store_embeddings(chunks, embeddings)
def test_empty_context():
assert generate_tests(diff="diff --git a/app.py b/app.py", context=[]) == ""

def test_query_embeddings():
query_embedding = [1, 2]
results = query_embeddings(query_embedding)
assert isinstance(results, list)
def test_intent_extraction_failure():
assert generate_tests(diff="diff --git a/app.py b/app.py", intent={"error": "intent extraction failed"}) == ""

def test_query_embeddings_empty():
collection = client.get_or_create_collection(name="empty")
query_embedding = [1, 2]
results = collection.query(query_embeddings=[query_embedding], n_results=5)
assert results.get("documents", []) == [[]]
def test_commit_tests_failure():
try:
commit_tests()
assert False
except Exception as e:
assert str(e) != ""

def test_query_embeddings_min_length():
query_embedding = [1, 2]
chunks = ["a" * 29, "b" * 31]
embeddings = [[1, 2], [3, 4]]
store_embeddings(chunks, embeddings)
results = query_embeddings(query_embedding)
assert len(results) == 1

def test_query_embeddings_k():
query_embedding = [1, 2]
chunks = ["a" * 31, "b" * 31, "c" * 31]
embeddings = [[1, 2], [3, 4], [5, 6]]
store_embeddings(chunks, embeddings)
results = query_embeddings(query_embedding, k=2)
assert len(results) == 2
def test_generate_review_with_intent():
diff = "diff --git a/app.py b/app.py"
context = ["def safe_divide(a, b):"]
intent = {"purpose": "generate tests", "properties": ["a", "b"], "edge_cases": ["a=0", "b=0"]}
assert generate_tests(diff, context, intent) != ""