ashkanhsn/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 this 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 125Respond with ### headers: Explanation, Fixed Code, Why the fix works.126 127### output:128### Fixed Code:129```python130"""131 temp = temperature_override if temperature_override is not None else 0.6132 max_tokens = 1024133 do_sample = True134 135 elif task == "chat":136 prompt = f"""### instruction:137{code}138 139### output:140"""141 temp = temperature_override if temperature_override is not None else 0.2142 max_tokens = 200143 do_sample = True144 145 else:146 return generate_response(code, "fix", use_base_model, temperature_override)147 148 # use ONE model object (PEFT) and toggle adapter off for base inference149 model_name = 'BASE' if use_base_model else 'FINE-TUNED'150 print(f'Using {model_name} model')151 152 inputs = tokenizer(prompt, return_tensors='pt')153 inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}154 155 # generate response156 ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()157 with torch.no_grad():158 with ctx:159 if do_sample:160 outputs = finetuned_model.generate(161 **inputs,162 max_new_tokens=max_tokens,163 temperature=temp,164 do_sample=True,165 top_p=0.95,166 pad_token_id=tokenizer.eos_token_id,167 eos_token_id=tokenizer.eos_token_id168 )169 else:170 outputs = finetuned_model.generate(171 **inputs,172 max_new_tokens=max_tokens,173 do_sample=False,174 pad_token_id=tokenizer.eos_token_id,175 eos_token_id=tokenizer.eos_token_id176 )177 178 response = tokenizer.decode(outputs[0], skip_special_tokens=True)179 if "### output:" in response:180 response = response.split("### output:")[-1].strip()181 182 return response183 184# Regex pattern used in "fix" task to split code from tests185TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")186 187def longest_valid_python_prefix(code: str) -> str:188 """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""189 lines = code.splitlines()190 best = ""191 buf = []192 for line in lines:193 buf.append(line)194 candidate = "\n".join(buf).strip() + "\n"195 try:196 ast.parse(candidate)197 best = "\n".join(buf)198 except SyntaxError:199 pass200 return best.strip() if best.strip() else code.strip()201 202 203def parse_response(text: str, task: str) -> CodeResponse:204 """Parse model response based on actual output format"""205 raw = text.strip()206 207 if task == "fix":208 # extract only the function/class definition209 code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()210 fixed_code = longest_valid_python_prefix(code_part)211 212 # Remove trailing comments and empty lines213 lines = fixed_code.split('\n')214 while lines:215 last_line = lines[-1].strip()216 if not last_line or last_line.startswith('#'):217 lines.pop()218 else:219 break220 221 fixed_code = '\n'.join(lines)222 223 return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)224 225 if task == "analyze":226 lines = raw.split('\n')227 228 # fixed code229 fixed_code = None230 code_lines = [] 231 232 fixed_code_match = re.search(233 r'###?\s*fixed\s+code.*?:?\s*\n+```(?:python)?\s*\n(.*?)\n```',234 raw,235 re.IGNORECASE | re.DOTALL236 )237 if fixed_code_match:238 fixed_code = fixed_code_match.group(1).strip()239 240 # Fallback: extract function/class definition from raw text241 if not fixed_code:242 in_function = False243 244 for line in lines:245 if line.startswith('def ') or line.startswith('class '):246 in_function = True247 code_lines.append(line)248 elif in_function and (line.startswith(' ') or line.startswith('\t') or 249 line.strip().startswith('"""') or line.strip().startswith("'''")):250 code_lines.append(line)251 elif in_function and line.strip() and not line.startswith(' '):252 break253 elif in_function:254 code_lines.append(line)255 256 fixed_code = '\n'.join(code_lines).strip() if code_lines else None257 258 # extract explanation259 explanation = None260 261 # "# Explanation:" comment section262 # captures everything from "# Explanation:" until "# Test" or end263 comment_explanation_match = re.search(264 r'#\s*Explanation[^:\n]*:?\s*\n((?:#[^\n]*\n?)+)',265 raw,266 re.IGNORECASE267 )268 if comment_explanation_match:269 explanation_text = comment_explanation_match.group(1)270 # remove # prefixes and join lines271 cleaned_lines = []272 for line in explanation_text.split('\n'):273 # remove leading # and whitespace274 clean = line.strip().lstrip('#').strip()275 # stop if we hit test-related content276 if clean.lower().startswith('test'):277 break278 if clean:279 cleaned_lines.append(clean)280 explanation = ' '.join(cleaned_lines) if cleaned_lines else None281 282 # Try markdown style ### Explanation:283 if not explanation:284 markdown_explanation_match = re.search(285 r'###?\s*(?:bug\s+)?explanation[^\n]*:?\s*\n+(.*?)(?=\n###?\s*(?:test|why|fixed)|#\s*Test|$)',286 raw,287 re.IGNORECASE | re.DOTALL288 )289 if markdown_explanation_match:290 explanation = markdown_explanation_match.group(1).strip()291 # Clean up any # prefixes292 cleaned_lines = []293 for line in explanation.split('\n'):294 clean = line.strip().lstrip('#').strip()295 if clean:296 cleaned_lines.append(clean)297 explanation = ' '.join(cleaned_lines) if cleaned_lines else None298 299 # extract test assertion300 test_lines = []301 302 test_start_match = re.search(303 r'(?:#\s*Test(?:s|ing)?[^\n]*|###?\s*Test(?:s|ing)?[^\n]*)\n',304 raw,305 re.IGNORECASE306 )307 308 if test_start_match:309 test_section = raw[test_start_match.end():]310 311 # extract all assert statements312 for line in test_section.split('\n'):313 stripped = line.strip()314 if stripped.startswith('assert '):315 # remove inline comments for cleaner output316 clean_assert = re.sub(r'\s*#.*$', '', stripped)317 test_lines.append(clean_assert)318 319 # Fallback: just find all assert statements in the raw text320 if not test_lines:321 for line in lines:322 stripped = line.strip()323 if stripped.startswith('assert '):324 clean_assert = re.sub(r'\s*#.*$', '', stripped)325 test_lines.append(clean_assert)326 327 tests = '\n'.join(test_lines) if test_lines else None328 329 # extract "WHY THE FIX WORKS" (mainly for base model)330 whyFix = None331 332 # Comment style: # Why the fix works:333 why_comment_match = re.search(334 r'#\s*Why\s+(?:the\s+)?fix\s+works:?\s*\n((?:#[^\n]*\n?)+)',335 raw,336 re.IGNORECASE337 )338 if why_comment_match:339 why_text = why_comment_match.group(1)340 cleaned_lines = []341 for line in why_text.split('\n'):342 clean = line.strip().lstrip('#').strip()343 if clean.lower().startswith('test'):344 break345 if clean:346 cleaned_lines.append(clean)347 whyFix = ' '.join(cleaned_lines) if cleaned_lines else None348 349 # Markdown style: ### Why the fix works:350 if not whyFix:351 why_markdown_match = re.search(352 r'###?\s*why\s+(?:the\s+)?fix\s+works[^\n]*:?\s*\n+(.*?)(?=\n###?\s*test|#\s*Test|Testing|$)',353 raw,354 re.IGNORECASE | re.DOTALL355 )356 if why_markdown_match:357 whyFix = why_markdown_match.group(1).strip()358 # Clean up359 cleaned_lines = []360 for line in whyFix.split('\n'):361 clean = line.strip().lstrip('#').strip()362 if clean:363 cleaned_lines.append(clean)364 whyFix = ' '.join(cleaned_lines) if cleaned_lines else None365 366 #bug summary367 detected = []368 if explanation:369 # Take first sentence for the summary370 sentences = explanation.split('.')371 for s in sentences:372 s = s.strip()373 if len(s) > 20:374 detected = [s + "."]375 break376 377 if not detected:378 detected = ["Potential bug detected - see explanation for details"]379 380 return CodeResponse(381 text=raw,382 detectedBugs=detected,383 fixedCode=fixed_code,384 explanation=explanation,385 whyFix=whyFix,386 tests=tests,387 )388 389 # For chat or unknown tasks390 return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)391 392 393@app.get("/")394def read_root():395 return {396 "name": "Python Bug Fixer API",397 "status": "running",398 "models": {399 "base": "microsoft/phi-2",400 "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"401 },402 "version": "2.3.0",403 "info": "Shared base model approach - efficient memory usage"404 }405 406@app.get("/health")407def health_check():408 return {409 "status": "healthy",410 "base_model_loaded": base_model is not None,411 "finetuned_model_loaded": finetuned_model is not None,412 "device": str(device)413 }414 415 416@app.post("/analyze", response_model=CodeResponse)417async def analyze_code(request: CodeRequest): 418 try:419 print(f"Request: task={request.task}, use_base_model={request.use_base_model}")420 response_text = generate_response(421 request.code, 422 request.task, 423 request.use_base_model,424 request.temperature425 )426 return parse_response(response_text, request.task)427 except Exception as e:428 print(f"Error: {str(e)}")429 raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")430 431if __name__ == "__main__":432 uvicorn.run(app, host="0.0.0.0", port=7860)