diff --git a/agent/llm/test_generator.py b/agent/llm/test_generator.py index 0701db4..dbd12b9 100644 --- a/agent/llm/test_generator.py +++ b/agent/llm/test_generator.py @@ -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", [])) diff --git a/agent/main.py b/agent/main.py index 33fb385..bde1d3b 100644 --- a/agent/main.py +++ b/agent/main.py @@ -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: @@ -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") diff --git a/tests/test_generated.py b/tests/test_generated.py index 48db62f..a7c68e4 100644 --- a/tests/test_generated.py +++ b/tests/test_generated.py @@ -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 \ No newline at end of file +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) != "" \ No newline at end of file