Spaces:
Running
Running
| 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 | |
| # ────────────────────────────────────────────── | |
| def health_check(): | |
| return { | |
| "status": "online", | |
| "model": "microsoft/codebert-base", | |
| "device": getattr(classifier, "device", "cpu"), | |
| "cached_repos": list(REPO_CACHE.keys()), | |
| } | |
| def get_usage(): | |
| return llm_engine.get_usage_stats() | |
| 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)) | |
| 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") | |
| 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") | |
| 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)) | |
| 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)) | |
| 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)) | |
| 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) | |
| # ────────────────────────────────────────────── | |
| 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)) | |
| 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") | |
| 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) |