Hannounaaa/phi2-python-bugfixer-api
0
1"""2HuggingFace Space - Python Bug Fixer API3FAST API backend for fine-tuned Phi-24 5FIXED VERSION: Better parsing of model output format6"""7 8import ast9import re10from contextlib import nullcontext11from typing import Optional, List, Dict12 13import torch14from fastapi import FastAPI, HTTPException15from fastapi.middleware.cors import CORSMiddleware16from pydantic import BaseModel, Field, ConfigDict17from transformers import AutoModelForCausalLM, AutoTokenizer18from peft import PeftModel19import uvicorn20 21 22app = FastAPI(23 title="Python Bug Fixer API",24 description="AI-powered Python bug detection and fixing",25 version="2.3.0" # Updated version26)27 28# Enable CORS for VS Code extension29app.add_middleware(30 CORSMiddleware,31 allow_origins=["*"],32 allow_credentials=True,33 allow_methods=["*"],34 allow_headers=["*"],35)36 37# Global variables for models38base_model = None39finetuned_model = None40tokenizer = None41device = None42 43try:44 device = "cuda" if torch.cuda.is_available() else "cpu"45 print(f"Using device: {device}")46 47 # tokenizer48 print("Load tokenizer")49 tokenizer = AutoTokenizer.from_pretrained(50 "microsoft/phi-2",51 trust_remote_code=True52 )53 tokenizer.pad_token = tokenizer.eos_token54 print("Tokenizer loaded!")55 56 # load one base model (will be shared)57 print("Loading Phi-2 base model (fp16)")58 base_model = AutoModelForCausalLM.from_pretrained(59 "microsoft/phi-2",60 trust_remote_code=True,61 torch_dtype=torch.float16 if device == "cuda" else torch.float32,62 device_map="auto" if device == "cuda" else None,63 )64 base_model.eval()65 print("Base model loaded!")66 67 # load PEFT adapter on top of base model68 print("Load fine-tuned PEFT adapter")69 finetuned_model = PeftModel.from_pretrained(70 base_model,71 "ashkanhsn/phi2-python-bugfixer"72 )73 finetuned_model.eval()74 print("Fine-tuned model ready!")75 76 print("Both models are ready!")77 78except Exception as e:79 print(f"Error loading models: {e}")80 import traceback81 traceback.print_exc()82 raise83 84 85class CodeRequest(BaseModel):86 model_config = ConfigDict(populate_by_name=True)87 code: str88 task: str = 'fix'89 temperature: Optional[float] = None90 # Accept both snake_case and camelCase from clients91 use_base_model: bool = Field(False, alias='useBaseModel')92 93 94class CodeResponse(BaseModel):95 text: str96 detectedBugs: list[str]97 fixedCode: Optional[str] = None98 explanation: Optional[str] = None99 whyFix: Optional[str] = None100 tests: Optional[str] = None101 error: Optional[str] = None102 103 104def generate_response(code: str, task: str, use_base_model: bool = False, temperature_override: Optional[float] = None) -> str:105 106 # Build prompt based on task107 if task == "fix":108 prompt = f"""### instruction:109Fix the following buggy Python code:110 111{code}112 113### output:114"""115 temp = temperature_override if temperature_override is not None else 0.2116 max_tokens = 300117 do_sample = temp > 0.0118 119 elif task == "analyze":120 prompt = f"""### instruction:121Fix this buggy Python code.122 123{code}124 125Provide:1261. Bug explanation1272. Fixed code1283. Why the fix works129 130### output:131"""132 temp = temperature_override if temperature_override is not None else 0.7133 max_tokens = 1024134 do_sample = True135 136 elif task == "chat":137 prompt = f"""### instruction:138{code}139 140### output:141"""142 temp = temperature_override if temperature_override is not None else 0.2143 max_tokens = 200144 do_sample = True145 146 else:147 return generate_response(code, "fix", use_base_model, temperature_override)148 149 # use ONE model object (PEFT) and toggle adapter off for base inference150 model_name = 'BASE' if use_base_model else 'FINE-TUNED'151 print(f'Using {model_name} model')152 153 inputs = tokenizer(prompt, return_tensors='pt')154 inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}155 156 # generate response157 ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()158 with torch.no_grad():159 with ctx:160 if do_sample:161 outputs = finetuned_model.generate(162 **inputs,163 max_new_tokens=max_tokens,164 temperature=temp,165 do_sample=True,166 top_p=0.95,167 pad_token_id=tokenizer.eos_token_id,168 eos_token_id=tokenizer.eos_token_id169 )170 else:171 outputs = finetuned_model.generate(172 **inputs,173 max_new_tokens=max_tokens,174 do_sample=False,175 pad_token_id=tokenizer.eos_token_id,176 eos_token_id=tokenizer.eos_token_id177 )178 179 response = tokenizer.decode(outputs[0], skip_special_tokens=True)180 if "### output:" in response:181 response = response.split("### output:")[-1].strip()182 183 return response184 185# Regex pattern used in "fix" task to split code from tests186TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")187 188def longest_valid_python_prefix(code: str) -> str:189 """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""190 lines = code.splitlines()191 best = ""192 buf = []193 for line in lines:194 buf.append(line)195 candidate = "\n".join(buf).strip() + "\n"196 try:197 ast.parse(candidate)198 best = "\n".join(buf)199 except SyntaxError:200 pass201 return best.strip() if best.strip() else code.strip()202 203 204def parse_response(text: str, task: str) -> CodeResponse:205 """Parse model response based on actual output format"""206 raw = text.strip()207 208 if task == "fix":209 # extract only the function/class definition210 code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()211 fixed_code = longest_valid_python_prefix(code_part)212 213 # Remove trailing comments and empty lines214 lines = fixed_code.split('\n')215 while lines:216 last_line = lines[-1].strip()217 if not last_line or last_line.startswith('#'):218 lines.pop()219 else:220 break221 222 fixed_code = '\n'.join(lines)223 224 return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)225 226 if task == "analyze":227 lines = raw.split('\n')228 229 # fixed code230 fixed_code = None231 code_lines = [] 232 233 fixed_code_match = re.search(234 r'###?\s*fixed\s+code.*?:?\s*\n+```(?:python)?\s*\n(.*?)\n```',235 raw,236 re.IGNORECASE | re.DOTALL237 )238 if fixed_code_match:239 fixed_code = fixed_code_match.group(1).strip()240 241 # Fallback: extract function/class definition from raw text242 if not fixed_code:243 in_function = False244 245 for line in lines:246 if line.startswith('def ') or line.startswith('class '):247 in_function = True248 code_lines.append(line)249 elif in_function and (line.startswith(' ') or line.startswith('\t') or 250 line.strip().startswith('"""') or line.strip().startswith("'''")):251 code_lines.append(line)252 elif in_function and line.strip() and not line.startswith(' '):253 break254 elif in_function:255 code_lines.append(line)256 257 fixed_code = '\n'.join(code_lines).strip() if code_lines else None258 259 # extract explanation260 explanation = None261 262 # "# Explanation:" comment section263 # captures everything from "# Explanation:" until "# Test" or end264 comment_explanation_match = re.search(265 r'#\s*Explanation[^:\n]*:?\s*\n((?:#[^\n]*\n?)+)',266 raw,267 re.IGNORECASE268 )269 if comment_explanation_match:270 explanation_text = comment_explanation_match.group(1)271 # remove # prefixes and join lines272 cleaned_lines = []273 for line in explanation_text.split('\n'):274 # remove leading # and whitespace275 clean = line.strip().lstrip('#').strip()276 # stop if we hit test-related content277 if clean.lower().startswith('test'):278 break279 if clean:280 cleaned_lines.append(clean)281 explanation = ' '.join(cleaned_lines) if cleaned_lines else None282 283 # Try markdown style ### Explanation:284 if not explanation:285 markdown_explanation_match = re.search(286 r'###?\s*(?:bug\s+)?explanation[^\n]*:?\s*\n+(.*?)(?=\n###?\s*(?:test|why|fixed)|#\s*Test|$)',287 raw,288 re.IGNORECASE | re.DOTALL289 )290 if markdown_explanation_match:291 explanation = markdown_explanation_match.group(1).strip()292 # Clean up any # prefixes293 cleaned_lines = []294 for line in explanation.split('\n'):295 clean = line.strip().lstrip('#').strip()296 if clean:297 cleaned_lines.append(clean)298 explanation = ' '.join(cleaned_lines) if cleaned_lines else None299 300 # extract test assertion301 test_lines = []302 303 test_start_match = re.search(304 r'(?:#\s*Test(?:s|ing)?[^\n]*|###?\s*Test(?:s|ing)?[^\n]*)\n',305 raw,306 re.IGNORECASE307 )308 309 if test_start_match:310 test_section = raw[test_start_match.end():]311 312 # extract all assert statements313 for line in test_section.split('\n'):314 stripped = line.strip()315 if stripped.startswith('assert '):316 # remove inline comments for cleaner output317 clean_assert = re.sub(r'\s*#.*$', '', stripped)318 test_lines.append(clean_assert)319 320 # Fallback: just find all assert statements in the raw text321 if not test_lines:322 for line in lines:323 stripped = line.strip()324 if stripped.startswith('assert '):325 clean_assert = re.sub(r'\s*#.*$', '', stripped)326 test_lines.append(clean_assert)327 328 tests = '\n'.join(test_lines) if test_lines else None329 330 # extract "WHY THE FIX WORKS" (mainly for base model)331 whyFix = None332 333 # Comment style: # Why the fix works:334 why_comment_match = re.search(335 r'#\s*Why\s+(?:the\s+)?fix\s+works:?\s*\n((?:#[^\n]*\n?)+)',336 raw,337 re.IGNORECASE338 )339 if why_comment_match:340 why_text = why_comment_match.group(1)341 cleaned_lines = []342 for line in why_text.split('\n'):343 clean = line.strip().lstrip('#').strip()344 if clean.lower().startswith('test'):345 break346 if clean:347 cleaned_lines.append(clean)348 whyFix = ' '.join(cleaned_lines) if cleaned_lines else None349 350 # Markdown style: ### Why the fix works:351 if not whyFix:352 why_markdown_match = re.search(353 r'###?\s*why\s+(?:the\s+)?fix\s+works[^\n]*:?\s*\n+(.*?)(?=\n###?\s*test|#\s*Test|Testing|$)',354 raw,355 re.IGNORECASE | re.DOTALL356 )357 if why_markdown_match:358 whyFix = why_markdown_match.group(1).strip()359 # Clean up360 cleaned_lines = []361 for line in whyFix.split('\n'):362 clean = line.strip().lstrip('#').strip()363 if clean:364 cleaned_lines.append(clean)365 whyFix = ' '.join(cleaned_lines) if cleaned_lines else None366 367 #bug summary368 detected = []369 if explanation:370 # Take first sentence for the summary371 sentences = explanation.split('.')372 for s in sentences:373 s = s.strip()374 if len(s) > 20:375 detected = [s + "."]376 break377 378 if not detected:379 detected = ["Potential bug detected - see explanation for details"]380 381 return CodeResponse(382 text=raw,383 detectedBugs=detected,384 fixedCode=fixed_code,385 explanation=explanation,386 whyFix=whyFix,387 tests=tests,388 )389 390 # For chat or unknown tasks391 return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)392 393 394@app.get("/")395def read_root():396 return {397 "name": "Python Bug Fixer API",398 "status": "running",399 "models": {400 "base": "microsoft/phi-2",401 "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"402 },403 "version": "2.3.0",404 "info": "Shared base model approach - efficient memory usage"405 }406 407@app.get("/health")408def health_check():409 return {410 "status": "healthy",411 "base_model_loaded": base_model is not None,412 "finetuned_model_loaded": finetuned_model is not None,413 "device": str(device)414 }415 416 417@app.post("/analyze", response_model=CodeResponse)418async def analyze_code(request: CodeRequest): 419 try:420 print(f"Request: task={request.task}, use_base_model={request.use_base_model}")421 response_text = generate_response(422 request.code, 423 request.task, 424 request.use_base_model,425 request.temperature426 )427 return parse_response(response_text, request.task)428 except Exception as e:429 print(f"Error: {str(e)}")430 raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")431 432if __name__ == "__main__":433 uvicorn.run(app, host="0.0.0.0", port=7860)