# FastAPI Backend for AI Code Security Scanner # Provides REST API for integration with other systems from fastapi import FastAPI, HTTPException, UploadFile, File, BackgroundTasks from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse from pydantic import BaseModel from typing import List, Optional, Dict, Any import uvicorn import json import asyncio from datetime import datetime import logging # Import our detectors from combined_detector import CombinedCodeDetector from rule_detector import RuleBasedCodeDetector # Setup logging logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', handlers=[ logging.FileHandler('api_logs.log'), logging.StreamHandler() ] ) logger = logging.getLogger(__name__) # Initialize FastAPI app app = FastAPI( title="AI Code Security Scanner API", description="REST API for detecting security vulnerabilities in Python code", version="1.0.0", docs_url="/docs", redoc_url="/redoc" ) # Add CORS middleware app.add_middleware( CORSMiddleware, allow_origins=["*"], # In production, restrict this allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # Initialize detectors detector = CombinedCodeDetector() rule_detector = RuleBasedCodeDetector() # Pydantic models for request/response class CodeRequest(BaseModel): code: str language: str = "python" analysis_mode: str = "combined" # "rules_only" or "combined" detailed: bool = False class CodeResponse(BaseModel): security_score: float issues_count: int issues: List[Dict[str, Any]] analysis_time: float timestamp: str model_used: str class BatchRequest(BaseModel): files: List[str] # List of code snippets analysis_mode: str = "combined" class BatchResponse(BaseModel): results: List[CodeResponse] summary: Dict[str, Any] class HealthResponse(BaseModel): status: str timestamp: str models_loaded: bool version: str # In-memory cache for recent requests request_cache = {} CACHE_SIZE = 100 def add_to_cache(key: str, result: dict): # Simple cache implementation if len(request_cache) >= CACHE_SIZE: # Remove oldest item oldest_key = next(iter(request_cache)) request_cache.pop(oldest_key) request_cache[key] = { "result": result, "timestamp": datetime.now().isoformat() } # Health check endpoint @app.get("/", response_model=HealthResponse) async def root(): # Root endpoint with health check return HealthResponse( status="healthy", timestamp=datetime.now().isoformat(), models_loaded=True, version="1.0.0" ) @app.get("/health", response_model=HealthResponse) async def health_check(): # Health check endpoint return HealthResponse( status="healthy", timestamp=datetime.now().isoformat(), models_loaded=True, version="1.0.0" ) # Single code analysis endpoint @app.post("/analyze", response_model=CodeResponse) async def analyze_code(request: CodeRequest): """ Analyze a single code snippet for security vulnerabilities - **code**: Python code to analyze - **language**: Programming language (default: python) - **analysis_mode**: "rules_only" or "combined" (default: combined) - **detailed**: Return detailed issue information """ start_time = datetime.now() # Check cache first cache_key = f"{request.code[:50]}_{request.analysis_mode}" if cache_key in request_cache: cached_result = request_cache[cache_key]["result"] logger.info(f"Serving from cache: {cache_key}") return JSONResponse(content=cached_result) try: if not request.code.strip(): raise HTTPException(status_code=400, detail="Empty code provided") logger.info(f"Analyzing code ({len(request.code)} chars), mode: {request.analysis_mode}") # Perform analysis if request.analysis_mode == "rules_only": result = rule_detector.analyze(request.code) model_used = "rule_based" else: result = detector.combined_analysis(request.code) model_used = "combined" # Calculate analysis time analysis_time = (datetime.now() - start_time).total_seconds() # Prepare response response_data = { "security_score": result["security_score"], "issues_count": result["issue_count"], "issues": result["issues"] if request.detailed else [], "analysis_time": analysis_time, "timestamp": datetime.now().isoformat(), "model_used": model_used, "summary": result.get("summary", {}), "ml_analysis": result.get("ml_analysis", {}) if request.detailed else {} } # Cache the result add_to_cache(cache_key, response_data) logger.info(f"Analysis complete. Score: {result['security_score']}, Issues: {result['issue_count']}") return CodeResponse(**response_data) except Exception as e: logger.error(f"Error analyzing code: {str(e)}") raise HTTPException(status_code=500, detail=f"Analysis error: {str(e)}") # Batch analysis endpoint @app.post("/analyze/batch", response_model=BatchResponse) async def analyze_batch(request: BatchRequest, background_tasks: BackgroundTasks): """ Analyze multiple code snippets in batch - **files**: List of code snippets - **analysis_mode**: "rules_only" or "combined" """ start_time = datetime.now() if not request.files: raise HTTPException(status_code=400, detail="No files provided") if len(request.files) > 100: raise HTTPException(status_code=400, detail="Maximum 100 files per batch") logger.info(f"Starting batch analysis of {len(request.files)} files") results = [] issues_summary = { "critical": 0, "high": 0, "medium": 0, "low": 0, "total_files": len(request.files) } # Process each file for i, code in enumerate(request.files): try: if request.analysis_mode == "rules_only": result = rule_detector.analyze(code) else: result = detector.combined_analysis(code) # Update summary if "summary" in result: for severity in ["critical", "high", "medium", "low"]: issues_summary[severity] += result["summary"].get(severity, 0) # Create response response = { "security_score": result["security_score"], "issues_count": result["issue_count"], "issues": result["issues"], "analysis_time": 0, # Would need individual timing "timestamp": datetime.now().isoformat(), "model_used": request.analysis_mode, "file_index": i } results.append(response) except Exception as e: logger.error(f"Error analyzing file {i}: {str(e)}") results.append({ "security_score": 0, "issues_count": 0, "issues": [{"type": "analysis_error", "message": str(e)}], "analysis_time": 0, "timestamp": datetime.now().isoformat(), "model_used": "error", "file_index": i }) total_time = (datetime.now() - start_time).total_seconds() # Calculate average score valid_scores = [r["security_score"] for r in results if r["security_score"] > 0] avg_score = sum(valid_scores) / len(valid_scores) if valid_scores else 0 issues_summary["average_security_score"] = avg_score issues_summary["total_analysis_time"] = total_time return BatchResponse( results=results, summary=issues_summary ) # File upload endpoint @app.post("/analyze/file") async def analyze_file(file: UploadFile = File(...)): # Analyze code from uploaded file if not file.filename.endswith('.py'): raise HTTPException(status_code=400, detail="Only .py files are supported") try: content = await file.read() code = content.decode('utf-8') # Use combined analysis for files result = detector.combined_analysis(code) return { "filename": file.filename, "security_score": result["security_score"], "issues_count": result["issue_count"], "critical_issues": result["summary"].get("critical", 0), "high_issues": result["summary"].get("high", 0), "analysis_time": datetime.now().isoformat() } except Exception as e: logger.error(f"Error processing file {file.filename}: {str(e)}") raise HTTPException(status_code=500, detail=f"File processing error: {str(e)}") # Statistics endpoint @app.get("/stats") async def get_statistics(): # Get API usage statistics return { "cache_size": len(request_cache), "cache_keys": list(request_cache.keys())[:5], "timestamp": datetime.now().isoformat(), "status": "operational" } # Clear cache endpoint (admin) @app.delete("/cache") async def clear_cache(): # Clear the request cache global request_cache cache_size = len(request_cache) request_cache = {} logger.info(f"Cache cleared. Removed {cache_size} entries.") return {"message": f"Cache cleared. Removed {cache_size} entries."} if __name__ == "__main__": logger.info("Starting FastAPI server...") uvicorn.run( "api_backend:app", host="0.0.0.0", port=8000, reload=True, log_level="info" )