gitgud-ai / app /main.py
CodeCommunity's picture
Update app/main.py
5b7d105 verified
Raw
History Blame Contribute Delete
12.5 kB
import os
import re
import logging
import traceback
import time
import asyncio
from typing import List, Optional, Dict, Any
from concurrent.futures import ThreadPoolExecutor
from dotenv import load_dotenv
from fastapi import FastAPI, HTTPException, status
from pydantic import BaseModel
import uvicorn
load_dotenv()
from app.predictor import classifier, guide_generator, reviewer
from app.core.model_loader import llm_engine
from app.services.danger_zone_service import DangerZoneService
from app.services.convention_service import convention_service
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
app = FastAPI(title="GitGud AI Service")
REPO_CACHE: Dict[str, Dict[str, List[float]]] = {}
executor = ThreadPoolExecutor(max_workers=10)
# ──────────────────────────────────────────────
# Pydantic Models
# ──────────────────────────────────────────────
class FileRequest(BaseModel):
fileName: str
content: Optional[str] = None
repoName: Optional[str] = None
class BatchReviewRequest(BaseModel):
files: List[FileRequest]
class GuideRequest(BaseModel):
repoName: str
filePaths: List[str]
class SearchRequest(BaseModel):
query: str
embeddings: Optional[Dict[str, List[float]]] = None
repoName: Optional[str] = None
class ChatRequest(BaseModel):
query: str
context: List[Dict[str, str]]
repoName: str
class DestructiveActionRequest(BaseModel):
repoOwner: str
repoName: str
action: str # "hard_reset" | "force_push"
targetRef: str # e.g. "origin/main" or a SHA
currentLocalCommits: Optional[List[str]] = None
class ReviewRequest(BaseModel):
files: List[FileRequest]
# ──────────────────────────────────────────────
# Helpers
# ──────────────────────────────────────────────
def calculate_repo_health(total_vulns: int, avg_maint: float) -> int:
base_score = avg_maint * 10.0
penalty = total_vulns * 8.0
return int(max(10.0, min(100.0, base_score - penalty)))
def sync_review_worker(file_list: List[FileRequest]):
logger.info(f"--- [DEBUG] Processing {len(file_list)} files for code review ---")
try:
results = reviewer.service.review_batch_code(file_list)
return results
except Exception as e:
logger.error(f"--- [ERROR] Exception during review processing: {e} ---", exc_info=True)
raise e
def parse_tree_to_list(raw_tree: str):
nodes = []
for line in raw_tree.strip().split('\n'):
if line.startswith("```") or not line.strip():
continue
level = line.count('|') + (line.count(' ') // 2)
name = re.sub(r'[|└├─]', '', line).strip()
name = re.sub(r'\[.*?\]', '', name).strip()
if name:
nodes.append({
"name": name,
"type": "file" if '.' in name else "folder",
"level": level
})
return nodes
# ──────────────────────────────────────────────
# Core Endpoints
# ──────────────────────────────────────────────
@app.get("/")
def health_check():
return {
"status": "online",
"model": "microsoft/codebert-base",
"device": getattr(classifier, "device", "cpu"),
"cached_repos": list(REPO_CACHE.keys()),
}
@app.get("/usage")
def get_usage():
return llm_engine.get_usage_stats()
@app.post("/classify")
async def classify_file(request: FileRequest):
try:
result = classifier.predict(request.fileName, request.content)
if request.repoName:
if request.repoName not in REPO_CACHE:
REPO_CACHE[request.repoName] = {}
REPO_CACHE[request.repoName][request.fileName] = result["embedding"]
return {
"fileName": request.fileName,
"layer": result["label"],
"confidence": result["confidence"],
"embedding": result["embedding"]
}
except Exception as e:
logger.error(f"Classify failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/review-batch-code")
async def review_batch_code(request: BatchReviewRequest):
try:
loop = asyncio.get_running_loop()
raw_results = await loop.run_in_executor(executor, sync_review_worker, request.files)
final_reviews = raw_results if isinstance(raw_results, list) else [raw_results]
return {"results": final_reviews}
except Exception as e:
logger.error(f"Batch review critical failure: {traceback.format_exc()}")
raise HTTPException(status_code=500, detail="Internal processing error")
@app.post("/repo-dashboard-stats")
async def get_dashboard_stats(request: BatchReviewRequest):
try:
loop = asyncio.get_running_loop()
raw_reviews = await loop.run_in_executor(executor, sync_review_worker, request.files)
if not isinstance(raw_reviews, list):
raw_reviews = [raw_reviews]
total_vulns = 0
maint_scores = []
found_apis = set()
api_regex = re.compile(r'(?:get|post|put|delete|patch)\([\'"]\/(.*?)[\'"]', re.IGNORECASE)
for i, current_review in enumerate(raw_reviews):
vulns = current_review.get("vulnerabilities", [])
total_vulns += len(vulns)
m_score = current_review.get("metrics", {}).get("maintainability", 8.0)
maint_scores.append(m_score)
content = request.files[i].content if i < len(request.files) else None
if content:
matches = api_regex.findall(content)
for match in matches:
found_apis.add(f"/{match}")
num_files = len(maint_scores)
avg_maint = (sum(maint_scores) / num_files) if num_files > 0 else 0.0
health_score = calculate_repo_health(total_vulns, avg_maint)
return {
"repo_health": health_score,
"health_label": "Excellent" if health_score > 85 else "Good" if health_score > 60 else "Critical",
"security_issues": total_vulns,
"performance_ratio": f"{int(avg_maint * 10)}%",
"exposed_apis": list(found_apis),
"total_files_processed": num_files,
"average_maintainability": round(avg_maint, 1)
}
except Exception as e:
logger.error(f"Stats failed: {e}")
raise HTTPException(status_code=500, detail="Failed to aggregate metrics")
@app.post("/analyze-file")
async def analyze_file(request: FileRequest):
try:
result = classifier.predict(request.fileName, request.content)
summary = classifier.generate_file_summary(request.content, request.fileName)
tags = classifier.extract_tags(request.content, request.fileName)
return {
"fileName": request.fileName,
"layer": result["label"],
"summary": summary,
"tags": tags,
"embedding": result["embedding"],
}
except Exception as e:
if "429" in str(e):
raise HTTPException(status_code=429, detail="Limit Reached")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/semantic-search")
async def semantic_search(request: SearchRequest):
try:
embeddings = request.embeddings
if not embeddings and request.repoName and request.repoName in REPO_CACHE:
embeddings = REPO_CACHE[request.repoName]
if not embeddings:
return {"results": []}
results = classifier.semantic_search(request.query, embeddings)
return {"results": results}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/chat")
async def chat(request: ChatRequest):
try:
context_str = ""
for item in request.context:
context_str += f"--- FILE: {item['fileName']} ---\n{item['content']}\n\n"
prompt = f"""You are "GitGud AI", an expert software architect.
Repository: "{request.repoName}"
CONTEXT: {context_str if request.context else "(NO CODE PROVIDED)"}
USER QUESTION: {request.query}"""
response = llm_engine.generate_text(prompt)
return {"response": response, "status": "success"}
except Exception as e:
if "429" in str(e):
return {"response": "⚠️ Daily limit reached. Try again in a bit!", "status": "quota_error"}
raise HTTPException(status_code=500, detail=str(e))
@app.post("/generate-guide")
async def generate_guide(request: GuideRequest):
try:
markdown = guide_generator.generate_markdown(request.repoName, request.filePaths)
tree_match = re.search(r"Project Structure\n\n(.*?)(?=\n\n|$)", markdown, re.S)
structured_tree = []
if tree_match:
structured_tree = parse_tree_to_list(tree_match.group(1))
return {
"markdown": markdown,
"structured_tree": structured_tree,
"project_name": request.repoName
}
except Exception as e:
if "429" in str(e):
raise HTTPException(status_code=429, detail="AI Quota Exceeded")
raise HTTPException(status_code=500, detail=str(e))
# ──────────────────────────────────────────────
# Sandbox / Danger Zone Endpoints (new)
# ──────────────────────────────────────────────
@app.post("/sandbox/simulate-destructive-action")
async def simulate_danger(request: DestructiveActionRequest):
try:
result = DangerZoneService.simulate_destructive_action(
owner=request.repoOwner,
repo_name=request.repoName,
action=request.action,
target_ref=request.targetRef,
current_local_shas=request.currentLocalCommits or [],
)
return result
except Exception as e:
logger.error(f"Danger zone simulation failed: {e}", exc_info=True)
raise HTTPException(status_code=500, detail=str(e))
@app.post("/sandbox/review-practice-pr")
async def review_practice(request: ReviewRequest):
"""
Reuses the exact same AIReviewerService that powers /review-batch-code.
"""
try:
loop = asyncio.get_running_loop()
results = await loop.run_in_executor(executor, sync_review_worker, request.files)
if not isinstance(results, list):
results = [results]
# One-sentence overall verdict
verdict_prompt = (
f"Give a single short sentence verdict on this practice PR. "
f"Be encouraging but honest. Results summary: {str(results)[:600]}"
)
try:
summary = llm_engine.generate(verdict_prompt, max_tokens=60)
except Exception:
summary = "Practice review completed."
return {
"results": results,
"isPractice": True,
"summary": summary.strip() if isinstance(summary, str) else "Practice review completed."
}
except Exception as e:
logger.error(f"Practice PR review failed: {e}", exc_info=True)
raise HTTPException(status_code=500, detail="Practice review failed")
@app.get("/sandbox/repo-conventions")
async def get_repo_conventions(owner: str, repo: str):
try:
return await convention_service.get_conventions(owner, repo)
except Exception as e:
logger.error(f"Conventions failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
# ──────────────────────────────────────────────
# Entry point
# ──────────────────────────────────────────────
if __name__ == "__main__":
port = int(os.environ.get("PORT", 7860))
uvicorn.run(app, host="0.0.0.0", port=port)