armaanalam commited on
Commit
93df7ed
·
verified ·
1 Parent(s): e9b3659

Upload 9 files

Browse files
Files changed (9) hide show
  1. code_analysis.py +53 -0
  2. documentation.py +61 -0
  3. embedding.py +30 -0
  4. file_creator.py +58 -0
  5. rag_chain.py +25 -0
  6. repository_loader.py +152 -0
  7. retriever.py +59 -0
  8. routes.py +71 -0
  9. splitter.py +61 -0
code_analysis.py ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import ast
2
+ from llm import get_llm_client, build_prompt
3
+ from rag.rag_chain import build_context
4
+ from rag.retriever import retrieve_relevant_chunks
5
+
6
+
7
+ def explain_function(file_path: str, function_name: str) -> str:
8
+ function_code = extract_function_source(file_path, function_name)
9
+ related_chunks = retrieve_relevant_chunks(f"usages of {function_name}", k=3)
10
+ context = build_context(related_chunks)
11
+
12
+ prompt = build_prompt(function_code, context, task_type="qa")
13
+ return get_llm_client().generate(prompt)
14
+
15
+
16
+ def extract_function_source(file_path: str, function_name: str) -> str:
17
+ with open(file_path, "r", encoding="utf-8") as f:
18
+ source = f.read()
19
+
20
+ tree = ast.parse(source)
21
+ for node in ast.walk(tree):
22
+ if isinstance(node, ast.FunctionDef) and node.name == function_name:
23
+ return ast.get_source_segment(source, node)
24
+
25
+ raise ValueError(f"Function '{function_name}' not found in {file_path}")
26
+
27
+
28
+ def detect_bugs(file_path: str) -> list[dict]:
29
+ with open(file_path, "r", encoding="utf-8") as f:
30
+ code = f.read()
31
+
32
+ prompt = build_prompt("", code, task_type="bug_finding")
33
+ raw_response = get_llm_client().generate(prompt)
34
+
35
+ import json
36
+ try:
37
+ return json.loads(raw_response)
38
+ except json.JSONDecodeError:
39
+ return [{"line": 0, "issue": "Could not parse model output", "severity": "unknown", "suggestion": raw_response}]
40
+
41
+
42
+ def analyze_complexity(file_path: str) -> dict:
43
+ from radon.complexity import cc_visit, cc_rank
44
+ with open(file_path, "r", encoding="utf-8") as f:
45
+ code = f.read()
46
+
47
+ results = cc_visit(code)
48
+ return {
49
+ "functions": [
50
+ {"name": r.name, "complexity": r.complexity, "rank": cc_rank(r.complexity)}
51
+ for r in results
52
+ ]
53
+ }
documentation.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from llm import get_llm_client, build_prompt
3
+ from rag.repository_loader import load_repository
4
+
5
+
6
+ def generate_docstring(function_code: str) -> str:
7
+ prompt = build_prompt(function_code, "", task_type="docstring")
8
+ return get_llm_client().generate(prompt)
9
+
10
+
11
+ def generate_module_docs(file_path: str) -> str:
12
+ with open(file_path, "r", encoding="utf-8") as f:
13
+ code = f.read()
14
+
15
+ prompt = build_prompt(
16
+ f"Generate documentation for this module: {file_path}",
17
+ code,
18
+ task_type="qa",
19
+ )
20
+ return get_llm_client().generate(prompt)
21
+
22
+
23
+ def summarize_repo_structure(root_path: str) -> dict:
24
+ documents = load_repository(root_path)
25
+
26
+ tech_stack = detect_tech_stack(root_path)
27
+ file_tree = build_file_tree(documents)
28
+
29
+ return {
30
+ "total_files": len(documents),
31
+ "tech_stack": tech_stack,
32
+ "file_tree": file_tree,
33
+ }
34
+
35
+
36
+ def detect_tech_stack(root_path: str) -> list[str]:
37
+ markers = {
38
+ "requirements.txt": "Python",
39
+ "package.json": "Node.js",
40
+ "go.mod": "Go",
41
+ "pom.xml": "Java (Maven)",
42
+ }
43
+ found = []
44
+ for marker_file, tech_name in markers.items():
45
+ if os.path.exists(os.path.join(root_path, marker_file)):
46
+ found.append(tech_name)
47
+ return found
48
+
49
+
50
+ def build_file_tree(documents: list) -> list[str]:
51
+ return sorted(doc.file_path for doc in documents)
52
+
53
+
54
+ def generate_readme(root_path: str) -> str:
55
+ summary = summarize_repo_structure(root_path)
56
+ prompt = build_prompt(
57
+ f"Generate a README.md for a project with this structure: {summary}",
58
+ "",
59
+ task_type="qa",
60
+ )
61
+ return get_llm_client().generate(prompt)
embedding.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pyrefly: ignore [missing-import]
2
+ from langchain_google_genai import GoogleGenerativeAIEmbeddings
3
+ from config import get_settings
4
+ from rag.repository_loader import CodeDocument
5
+
6
+
7
+ def get_embedding_client() -> GoogleGenerativeAIEmbeddings:
8
+ settings = get_settings()
9
+
10
+ return GoogleGenerativeAIEmbeddings(
11
+ model=settings.embedding_model,
12
+ google_api_key=settings.gemini_api_key,
13
+ )
14
+
15
+
16
+ def embed_document(chunks: list[dict]) -> list[dict]:
17
+ embeddings = get_embedding_client()
18
+ texts = [chunk["content"] for chunk in chunks]
19
+ vectors = embeddings.embed_documents(texts)
20
+
21
+ for chunk, vector in zip(chunks, vectors):
22
+ chunk["embedding"] = vector
23
+ return chunks
24
+
25
+
26
+ def embed_query(query: str) -> list[float]:
27
+ embeddings = get_embedding_client()
28
+ return embeddings.embed_query(query)
29
+
30
+
file_creator.py ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from dataclasses import dataclass, asdict
2
+ from llm import get_llm_client, build_prompt
3
+ from rag.rag_chain import build_context
4
+ from rag.retriever import retrieve_relevant_chunks
5
+
6
+
7
+ @dataclass
8
+ class FileProposal:
9
+ path: str
10
+ content: str
11
+ reason: str
12
+ approved: bool = False
13
+
14
+
15
+ def propose_new_file(description: str, context_query: str | None = None) -> FileProposal:
16
+ context = ""
17
+ if context_query:
18
+ related_chunks = retrieve_relevant_chunks(context_query, k=5)
19
+ context = build_context(related_chunks)
20
+
21
+ prompt = build_prompt(description, context, task_type="file_creation")
22
+ content = get_llm_client().generate(prompt)
23
+
24
+ proposed_path = infer_file_path(description)
25
+
26
+ return FileProposal(path=proposed_path, content=content, reason=description)
27
+
28
+
29
+ def infer_file_path(description: str) -> str:
30
+ prompt = f"Given this file creation request: '{description}', respond with ONLY a suitable relative file path, nothing else."
31
+ path = get_llm_client().generate(prompt).strip()
32
+ return path
33
+
34
+
35
+ async def apply_approved_file(proposal: FileProposal, user_confirmed: bool, base_dir: str = ".") -> dict:
36
+ """Write the proposed file to local disk under *base_dir*.
37
+
38
+ The LLM-suggested path (e.g. /src/geometry/rectangle.js) is treated as
39
+ relative to *base_dir*, so leading slashes/backslashes are stripped before
40
+ joining to avoid accidental absolute-path writes.
41
+ """
42
+ if not user_confirmed:
43
+ return {"status": "rejected", "proposal": asdict(proposal)}
44
+
45
+ import os
46
+
47
+ # Strip any leading separators so the path is always relative to base_dir
48
+ relative_path = proposal.path.lstrip("/\\")
49
+ abs_path = os.path.join(base_dir, relative_path)
50
+
51
+ # Create parent directories if they don't exist
52
+ os.makedirs(os.path.dirname(abs_path), exist_ok=True)
53
+
54
+ with open(abs_path, "w", encoding="utf-8") as f:
55
+ f.write(proposal.content)
56
+
57
+ proposal.path = abs_path # update so the caller can display the real path
58
+ return {"status": "written", "path": abs_path}
rag_chain.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from rag.retriever import retrieve_relevant_chunks
2
+ from llm import get_llm_client, build_prompt
3
+
4
+
5
+ def build_context(chunks: list[dict]) -> str:
6
+ parts = []
7
+ for c in chunks:
8
+ meta = c["metadata"]
9
+ header = f"# {meta['file_path']} (lines {meta['start_line']}-{meta['end_line']})"
10
+ parts.append(f"{header}\n{c['content']}")
11
+ return "\n\n---\n\n".join(parts)
12
+
13
+
14
+ def run_rag_query(query: str, k: int = 5) -> dict:
15
+ chunks = retrieve_relevant_chunks(query, k)
16
+ context = build_context(chunks)
17
+
18
+ prompt = build_prompt(query, context, task_type="qa")
19
+ llm = get_llm_client()
20
+ answer = llm.generate(prompt)
21
+
22
+ return {
23
+ "answer": answer,
24
+ "sources": [c["metadata"] for c in chunks],
25
+ }
repository_loader.py ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List
2
+ from config import get_settings
3
+ import os
4
+ from dataclasses import dataclass
5
+
6
+
7
+ @dataclass
8
+ class CodeDocument:
9
+ content: str
10
+ file_path: str
11
+ language: str
12
+ size_bytes: int
13
+
14
+
15
+ IGNORE_DIRS = {".git", "node_modules", "__pycache__", "venv", ".venv", "dist", "build"}
16
+ LANGUAGE_BY_EXT = {
17
+ ".py": "python",
18
+ ".js": "javascript",
19
+ ".ts": "typescript",
20
+ ".java": "java",
21
+ ".go": "go",
22
+ ".md": "markdown",
23
+ }
24
+
25
+
26
+ def load_repository(root_path: str) -> List[CodeDocument]:
27
+ documents: List[CodeDocument] = []
28
+ for dirpath, dirnames, filenames in os.walk(root_path):
29
+ dirnames[:] = [d for d in dirnames if d not in IGNORE_DIRS]
30
+ for filename in filenames:
31
+ fpath = os.path.join(dirpath, filename)
32
+ if should_include(fpath):
33
+ documents.append(read_file_with_metadata(fpath))
34
+ return documents
35
+
36
+
37
+ def should_include(fpath: str) -> bool:
38
+ settings = get_settings()
39
+ _, ext = os.path.splitext(fpath)
40
+ if ext not in settings.allowed_extensions:
41
+ return False
42
+ try:
43
+ size_kb = os.path.getsize(fpath)/1024
44
+ except OSError:
45
+ return False
46
+ return size_kb <= settings.max_file_size_kb
47
+
48
+
49
+ def read_file_with_metadata(filepath: str) -> CodeDocument:
50
+ with open(filepath, "r", encoding="utf-8", errors="ignore") as f:
51
+ content = f.read()
52
+ return CodeDocument(
53
+ content=content,
54
+ file_path=filepath,
55
+ language=detect_language(filepath),
56
+ size_bytes=len(content.encode("utf-8")),
57
+ )
58
+
59
+
60
+ def detect_language(filepath: str) -> str:
61
+ _, ext = os.path.splitext(filepath)
62
+ return LANGUAGE_BY_EXT.get(ext, "unknown")
63
+
64
+
65
+
66
+
67
+
68
+ """
69
+ import os
70
+ from dataclasses import dataclass
71
+ from langchain_community.document_loaders import DirectoryLoader, TextLoader
72
+ from config import get_settings
73
+
74
+
75
+ @dataclass
76
+ class CodeDocument:
77
+ content: str
78
+ file_path: str
79
+ language: str
80
+ size_bytes: int
81
+
82
+ LANGUAGE_BY_EXT = {
83
+ ".py": "python",
84
+ ".js": "javascript",
85
+ ".ts": "typescript",
86
+ ".java": "java",
87
+ ".go": "go",
88
+ ".cpp": "cpp",
89
+ ".c": "c",
90
+ ".cs": "csharp",
91
+ ".html": "html",
92
+ ".css": "css",
93
+ ".json": "json",
94
+ ".yaml": "yaml",
95
+ ".yml": "yaml",
96
+ ".md": "markdown",
97
+ }
98
+ def detect_language(filepath: str) -> str:
99
+ _, ext = os.path.splitext(filepath)
100
+ return LANGUAGE_BY_EXT.get(ext.lower(), "unknown")
101
+
102
+ def load_repository(root_path: str) -> list[CodeDocument]:
103
+ settings = get_settings()
104
+ loader = DirectoryLoader(
105
+ path=root_path,
106
+ glob="**/*",
107
+ recursive=True,
108
+ silent_errors=True,
109
+ loader_cls=TextLoader,
110
+ loader_kwargs={
111
+ "encoding": "utf-8",
112
+ "autodetect_encoding": True,
113
+ },
114
+ exclude=[
115
+ "**/.git/**",
116
+ "**/.venv/**",
117
+ "**/venv/**",
118
+ "**/__pycache__/**",
119
+ "**/node_modules/**",
120
+ "**/dist/**",
121
+ "**/build/**",
122
+ ],
123
+ )
124
+ documents = loader.load()
125
+ code_documents: list[CodeDocument] = []
126
+
127
+ for doc in documents:
128
+ filepath = doc.metadata["source"]
129
+ _, ext = os.path.splitext(filepath)
130
+
131
+ if ext.lower() not in settings.allowed_extensions:
132
+ continue
133
+
134
+ try:
135
+ size_bytes = os.path.getsize(filepath)
136
+ except OSError:
137
+ continue
138
+
139
+ if size_bytes > settings.max_file_size_kb * 1024:
140
+ continue
141
+
142
+ code_documents.append(
143
+ CodeDocument(
144
+ content=doc.page_content,
145
+ file_path=filepath,
146
+ language=detect_language(filepath),
147
+ size_bytes=size_bytes,
148
+ )
149
+ )
150
+
151
+ return code_documents
152
+ """
retriever.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import chromadb
2
+ from config import get_settings
3
+ from rag.embedding import embed_query
4
+
5
+
6
+ def get_chroma_client():
7
+ settings = get_settings()
8
+ return chromadb.PersistentClient(path=settings.vector_db_path)
9
+
10
+
11
+ def build_vector_store(embedded_chunks: list[dict]):
12
+ client = get_chroma_client()
13
+
14
+ # Always start fresh — delete any data from the previously ingested repo
15
+ # so that stale chunks never bleed into the current analysis session.
16
+ try:
17
+ client.delete_collection("codebase")
18
+ except Exception:
19
+ pass # collection didn't exist yet — that's fine
20
+
21
+ collection = client.create_collection("codebase")
22
+
23
+ ids = [f"{c['file_path']}::{c['chunk_index']}" for c in embedded_chunks]
24
+ embeddings = [c["embedding"] for c in embedded_chunks]
25
+ documents = [c["content"] for c in embedded_chunks]
26
+ metadatas = [
27
+ {
28
+ "file_path": c["file_path"],
29
+ "language": c["language"],
30
+ "start_line": c["start_line"],
31
+ "end_line": c["end_line"],
32
+ }
33
+ for c in embedded_chunks
34
+ ]
35
+ collection.add(ids=ids, embeddings=embeddings, documents=documents, metadatas=metadatas)
36
+ return collection
37
+
38
+
39
+ def load_vector_store():
40
+ client = get_chroma_client()
41
+ return client.get_or_create_collection("codebase")
42
+
43
+
44
+ def retrieve_relevant_chunks(query: str, k: int = 5) -> list[dict]:
45
+ collection = load_vector_store()
46
+ query_vector = embed_query(query)
47
+ results = collection.query(query_embeddings=[query_vector], n_results=k)
48
+ return format_results(results)
49
+
50
+
51
+ def format_results(results: dict) -> list[dict]:
52
+ formatted = []
53
+ documents = results.get("documents", [[]])[0]
54
+ metadatas = results.get("metadatas", [[]])[0]
55
+
56
+ for doc_text, meta in zip(documents, metadatas):
57
+ formatted.append({"content": doc_text, "metadata": meta})
58
+
59
+ return formatted
routes.py ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from fastapi import APIRouter
2
+ from pydantic import BaseModel
3
+
4
+ from rag.rag_chain import run_rag_query
5
+ from services.code_analysis import detect_bugs, explain_function, analyze_complexity
6
+ from services.documentation import generate_module_docs, generate_readme
7
+ from services.file_creator import propose_new_file, apply_approved_file, FileProposal
8
+
9
+ router = APIRouter()
10
+
11
+
12
+ class QueryRequest(BaseModel):
13
+ query: str
14
+ k: int = 5
15
+
16
+
17
+ class ExplainRequest(BaseModel):
18
+ file_path: str
19
+ function_name: str
20
+
21
+
22
+ class FileProposalRequest(BaseModel):
23
+ description: str
24
+ context_query: str | None = None
25
+
26
+
27
+ class FileApprovalRequest(BaseModel):
28
+ proposal: dict
29
+ approved: bool
30
+
31
+
32
+ @router.post("/query")
33
+ def query(req: QueryRequest):
34
+ return run_rag_query(req.query, req.k)
35
+
36
+
37
+ @router.post("/analyze/bugs")
38
+ def bugs(file_path: str):
39
+ return {"bugs": detect_bugs(file_path)}
40
+
41
+
42
+ @router.post("/analyze/complexity")
43
+ def complexity(file_path: str):
44
+ return analyze_complexity(file_path)
45
+
46
+
47
+ @router.post("/analyze/explain")
48
+ def explain(req: ExplainRequest):
49
+ return {"explanation": explain_function(req.file_path, req.function_name)}
50
+
51
+
52
+ @router.post("/docs/module")
53
+ def module_docs(file_path: str):
54
+ return {"docs": generate_module_docs(file_path)}
55
+
56
+
57
+ @router.post("/docs/readme")
58
+ def readme(root_path: str):
59
+ return {"readme": generate_readme(root_path)}
60
+
61
+
62
+ @router.post("/files/propose")
63
+ def propose(req: FileProposalRequest):
64
+ proposal = propose_new_file(req.description, req.context_query)
65
+ return proposal
66
+
67
+
68
+ @router.post("/files/approve")
69
+ async def approve(req: FileApprovalRequest):
70
+ proposal = FileProposal(**req.proposal)
71
+ return await apply_approved_file(proposal, req.approved)
splitter.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from langchain_text_splitters import RecursiveCharacterTextSplitter, Language
2
+ from rag.repository_loader import CodeDocument
3
+ from config import get_settings
4
+
5
+
6
+ LANGUAGE_MAP = {
7
+ "python": Language.PYTHON,
8
+ "javascript": Language.JS,
9
+ "typescript": Language.TS,
10
+ "java": Language.JAVA,
11
+ "go": Language.GO,
12
+ }
13
+
14
+
15
+ def split_code(doc : CodeDocument) -> list[dict]:
16
+
17
+ settings = get_settings()
18
+ lang_enum = LANGUAGE_MAP.get(doc.language.lower())
19
+
20
+ if lang_enum:
21
+ splitter = RecursiveCharacterTextSplitter.from_language(
22
+ language=lang_enum,
23
+ chunk_size=settings.chunk_size,
24
+ chunk_overlap=settings.chunk_overlap,
25
+ )
26
+ else:
27
+ splitter = RecursiveCharacterTextSplitter(
28
+ chunk_size=settings.chunk_size,
29
+ chunk_overlap=settings.chunk_overlap
30
+ )
31
+
32
+ raw_chunks = splitter.split_text(doc.content)
33
+ return [
34
+ attach_metadata(chunk, doc, idx)
35
+ for idx, chunk in enumerate(raw_chunks)
36
+ ]
37
+
38
+
39
+ def attach_metadata(chunk_text: str, doc: CodeDocument, index: int) -> dict:
40
+ start_line, end_line = estimate_line_range(doc.content, chunk_text)
41
+ return {
42
+ "content": chunk_text,
43
+ "file_path": doc.file_path,
44
+ "language": doc.language,
45
+ "chunk_index": index,
46
+ "start_line": start_line,
47
+ "end_line": end_line,
48
+ }
49
+
50
+
51
+ def estimate_line_range(full_text: str, chunk_text: str) -> tuple[int, int]:
52
+ offset = full_text.find(chunk_text)
53
+ if offset == -1:
54
+ return (0, 0)
55
+ start_line = full_text[:offset].count("\n") + 1
56
+ end_line = start_line + chunk_text.count("\n")
57
+ return (start_line, end_line)
58
+
59
+
60
+
61
+