diff --git a/agent/indexing/vector_store.py b/agent/indexing/vector_store.py index 41b615a..d180095 100644 --- a/agent/indexing/vector_store.py +++ b/agent/indexing/vector_store.py @@ -1,7 +1,9 @@ import chromadb +# safer collection creation client = chromadb.Client() -collection = client.create_collection(name="codebase") +collection = client.get_or_create_collection(name="codebase") + def store_embeddings(chunks, embeddings): for i, chunk in enumerate(chunks): @@ -11,23 +13,48 @@ def store_embeddings(chunks, embeddings): ids=[str(i)] ) + +MIN_CHUNK_LENGTH = 30 # Phase 4 threshold + + def query_embeddings(query_embedding, k=5): results = collection.query( query_embeddings=[query_embedding], n_results=k ) - docs = results.get("documents", [[]])[0] + docs = results.get("documents", [[]]) + + # safety check + if not docs or not docs[0]: + print("No documents retrieved") + return [] + + docs = docs[0] + + # ๐Ÿงน CLEAN: remove empty / noisy chunks + cleaned = [ + doc for doc in docs + if doc and len(doc.strip()) > MIN_CHUNK_LENGTH + ] + + # ๐ŸŽฏ PRIORITIZE: functions/classes + filtered = [ + doc for doc in cleaned + if "def " in doc or "class " in doc + ] + + # fallback if nothing found + if not filtered: + filtered = cleaned - # remove empty / tiny chunks - cleaned = [doc for doc in docs if doc and len(doc.strip()) > 20] + # ๐Ÿ” REMOVE DUPLICATES + unique_docs = list(dict.fromkeys(filtered)) - # ๐Ÿ”ฅ prioritize useful code (functions/classes) - filtered = [] - for doc in cleaned: - if "def " in doc or "class " in doc: - filtered.append(doc) + # ๐Ÿงช DEBUG LOGS + print(f"Retrieved {len(unique_docs)} relevant chunks") - print(f"Retrieved {len(filtered)} relevant chunks") + for i, doc in enumerate(unique_docs[:2]): + print(f"Chunk {i+1} preview:", doc[:100]) - return filtered[:k] \ No newline at end of file + return unique_docs[:k] \ No newline at end of file diff --git a/tests/test_generated.py b/tests/test_generated.py index ff04abe..ee02b1e 100644 --- a/tests/test_generated.py +++ b/tests/test_generated.py @@ -1,26 +1,34 @@ import pytest -import os -from agent.github.committer import commit_tests -from agent.main import get_pr_diff +from your_module import store_embeddings, query_embeddings -def test_commit_tests(): - commit_tests() +def test_store_embeddings(): + chunks = ["chunk1", "chunk2"] + embeddings = [[1, 2], [3, 4]] + store_embeddings(chunks, embeddings) -def test_get_pr_diff_empty(): - with pytest.raises(subprocess.CalledProcessError): - get_pr_diff() +def test_query_embeddings(): + query_embedding = [1, 2] + results = query_embeddings(query_embedding) + assert isinstance(results, list) -def test_commit_tests_exception(): - try: - commit_tests() - except Exception as e: - assert str(e) +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_get_pr_diff_no_diff(): - diff = get_pr_diff() - assert diff.strip() == "" +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_commit_tests_push(): - subprocess.run(["git", "config", "user.name", "github-actions"], check=True) - subprocess.run(["git", "config", "user.email", "actions@github.com"], check=True) - commit_tests() \ No newline at end of file +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