Sarra22/phi2-python-bugfixer-api
0
1"""2HuggingFace Space - Python Bug Fixer API3FAST API backend for fine-tuned Phi-24"""5 6import ast7import re8from contextlib import nullcontext9from typing import Optional, List, Dict10 11import torch12from fastapi import FastAPI, HTTPException13from fastapi.middleware.cors import CORSMiddleware14from pydantic import BaseModel, Field, ConfigDict15from transformers import AutoModelForCausalLM, AutoTokenizer16from peft import PeftModel17import uvicorn18 19 20app = FastAPI(21 title="Python Bug Fixer API",22 description="AI-powered Python bug detection and fixing",23 version="2.2.0"24)25 26# Enable CORS for VS Code extension27app.add_middleware(28 CORSMiddleware,29 allow_origins=["*"],30 allow_credentials=True,31 allow_methods=["*"],32 allow_headers=["*"],33)34 35# Global variables for models36base_model = None37finetuned_model = None38tokenizer = None39device = None40 41try:42 device = "cuda" if torch.cuda.is_available() else "cpu"43 print(f"Using device: {device}")44 45 # tokenizer46 print("Load tokenizer")47 tokenizer = AutoTokenizer.from_pretrained(48 "microsoft/phi-2",49 trust_remote_code=True50 )51 tokenizer.pad_token = tokenizer.eos_token52 print("Tokenizer loaded!")53 54 # load one base model (will be shared)55 print("Loading Phi-2 base model (fp16)")56 base_model = AutoModelForCausalLM.from_pretrained(57 "microsoft/phi-2",58 trust_remote_code=True,59 torch_dtype=torch.float16 if device == "cuda" else torch.float32,60 device_map="auto" if device == "cuda" else None,61 )62 base_model.eval()63 print("Base model loaded!")64 65 # load PEFT adapter on top of base model66 print("Load fine-tuned PEFT adapter")67 finetuned_model = PeftModel.from_pretrained(68 base_model,69 "ashkanhsn/phi2-python-bugfixer"70 )71 finetuned_model.eval()72 print("Fine-tuned model ready!")73 74 print("Both models are ready!")75 76except Exception as e:77 print(f"Error loading models: {e}")78 import traceback79 traceback.print_exc()80 raise81 82 83class CodeRequest(BaseModel):84 model_config = ConfigDict(populate_by_name=True)85 code: str86 task: str = 'fix'87 temperature: Optional[float] = None88 # Accept both snake_case and camelCase from clients89 use_base_model: bool = Field(False, alias='useBaseModel')90 91 92class CodeResponse(BaseModel):93 text: str94 detectedBugs: list[str]95 fixedCode: Optional[str] = None96 explanation: Optional[str] = None97 whyFix: Optional[str] = None98 tests: Optional[str] = None99 error: Optional[str] = None100 101 102def generate_response(code: str, task: str, use_base_model: bool = False, temperature_override: Optional[float] = None) -> str:103 104 # Build prompt based on task105 if task == "fix":106 prompt = f"""### instruction:107Fix the following buggy Python code:108 109{code}110 111### output:112"""113 temp = temperature_override if temperature_override is not None else 0.2114 max_tokens = 300115 do_sample = temp > 0.0116 117 elif task == "analyze":118 prompt = f"""### instruction:119Fix this buggy Python code.120 121{code}122 123Provide:1241. Bug explanation1252. Fixed code1263. Why the fix works127 128### output:129"""130 temp = temperature_override if temperature_override is not None else 0.7131 max_tokens = 1024132 do_sample = True133 134 elif task == "chat":135 prompt = f"""### instruction:136{code}137 138### output:139"""140 temp = temperature_override if temperature_override is not None else 0.2141 max_tokens = 200142 do_sample = True143 144 else:145 return generate_response(code, "fix", use_base_model, temperature_override)146 147 # use ONE model object (PEFT) and toggle adapter off for base inference148 model_name = 'BASE' if use_base_model else 'FINE-TUNED'149 print(f'Using {model_name} model')150 151 inputs = tokenizer(prompt, return_tensors='pt')152 inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}153 154 # generate response155 ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()156 with torch.no_grad():157 with ctx:158 if do_sample:159 outputs = finetuned_model.generate(160 **inputs,161 max_new_tokens=max_tokens,162 temperature=temp,163 do_sample=True,164 top_p=0.95,165 pad_token_id=tokenizer.eos_token_id,166 eos_token_id=tokenizer.eos_token_id167 )168 else:169 outputs = finetuned_model.generate(170 **inputs,171 max_new_tokens=max_tokens,172 do_sample=False,173 pad_token_id=tokenizer.eos_token_id,174 eos_token_id=tokenizer.eos_token_id175 )176 177 response = tokenizer.decode(outputs[0], skip_special_tokens=True)178 if "### output:" in response:179 response = response.split("### output:")[-1].strip()180 181 return response182 183STOP_HEADING_RE = re.compile(184 r"""(?im)^\s*(?:###\s*(?:explanation|why.*works|testing|tests?)\s*:?185 |#\s*Explanation(?:\s+of\s+the\s+fix)?\s*:?186 |Explanation\s+of\s+the\s+fix\s*:?187 )\s*$""",188 re.VERBOSE,189)190TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")191TEST_START_RE = re.compile(r"(?m)^(assert\s+|print\()")192FENCED_CODE_RE = re.compile(r"```(?:python)?\s*\n([\s\S]*?)\n```", re.IGNORECASE)193 194def extract_first_fenced_code(text: str) -> Optional[str]:195 m = FENCED_CODE_RE.search(text)196 return m.group(1).strip() if m else None197 198def longest_valid_python_prefix(code: str) -> str:199 """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""200 lines = code.splitlines()201 best = ""202 buf = []203 for line in lines:204 buf.append(line)205 candidate = "\n".join(buf).strip() + "\n"206 try:207 ast.parse(candidate)208 best = "\n".join(buf)209 except SyntaxError:210 pass211 return best.strip() if best.strip() else code.strip()212 213def extract_section(text: str, heading_patterns: List[str]) -> Optional[str]:214 """extract section content following a heading, stopping at next heading or top-level tests"""215 for hp in heading_patterns:216 m = re.search(hp, text, flags=re.IGNORECASE | re.MULTILINE)217 if not m:218 continue219 start = m.end()220 tail = text[start:]221 222 # stop at next heading or tests223 stop_m = re.search(r"(?im)^\s*###\s*[A-Za-z].*?:?\s*$|(?m)^(assert\s+|print\()", tail)224 chunk = tail[: stop_m.start()] if stop_m else tail225 chunk = chunk.strip()226 if len(chunk) >= 5:227 return chunk228 return None229 230def extract_tests(text: str) -> Optional[str]:231 # prefer explicit "Tests" section232 t = extract_section(text, [233 r"(?im)^\s*###\s*Tests?\s*:?\s*$",234 r"(?im)^\s*###\s*Testing\s*:?\s*$",235 r"(?im)^\s*#\s*Tests?\s*:?\s*$",236 r"(?im)^\s*#\s*Testing\s*:?\s*$",237 ])238 if t:239 # keep only assert/print lines if the section contains mixed content240 lines = [ln.rstrip() for ln in t.splitlines() if ln.strip()]241 # If there are asserts, show from first assert onward242 for i, ln in enumerate(lines):243 if ln.lstrip().startswith("assert"):244 return "\n".join(lines[i:]).strip()245 return "\n".join(lines).strip()246 247 # Fallback: from first top-level assert to end248 m = re.search(r"(?m)^assert\s+.*$", text)249 if m:250 return text[m.start():].strip()251 return None252 253def guess_bug_summaries(explanation: Optional[str]) -> List[str]:254 # if not explanation:255 # return ["Code analyzed and fixed - check explanation for details"]256 # # Take first 1-2 decent sentences257 # cleaned = " ".join([ln.strip().lstrip("#").strip() for ln in explanation.splitlines() if ln.strip()])258 # sentences = re.split(r"[.!?]\s+", cleaned)259 # out = []260 # for s in sentences:261 # s = s.strip()262 # if len(s) >= 18:263 # out.append(s + ".")264 # if len(out) == 2:265 # break266 # return out if out else ["Code analyzed and fixed - check explanation for details"]267 return []268 269def parse_response(text: str, task: str) -> CodeResponse:270 """parse model response"""271 raw = text.strip()272 273 if task == "fix":274 # extract only the function/class definition275 code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()276 fixed_code = longest_valid_python_prefix(code_part)277 return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)278 279 if task == "analyze":280 lines = raw.split('\n')281 282 # extract the fixed code (function/class definition only)283 code_lines = []284 in_function = False285 286 for line in lines:287 if line.startswith('def ') or line.startswith('class '):288 in_function = True289 code_lines.append(line)290 elif in_function and (line.startswith(' ') or line.startswith('\t') or line.strip().startswith('"""') or line.strip().startswith("'''")):291 code_lines.append(line)292 elif in_function and line.strip() and not line.startswith(' '):293 break294 elif in_function:295 code_lines.append(line)296 297 fixed_code = '\n'.join(code_lines).strip()298 299 # extract explanatory comments (lines starting with #)300 explanation_lines = []301 for line in lines:302 stripped = line.strip()303 # Get comment lines that are NOT part of the function definition304 if stripped.startswith('#') and line not in code_lines:305 explanation_lines.append(stripped.lstrip('#').strip())306 307 explanation = '\n'.join(explanation_lines).strip() if explanation_lines else None308 309 # extract test assertions310 test_lines = []311 for line in lines:312 stripped = line.strip()313 if stripped.startswith('assert '):314 test_lines.append(stripped)315 316 tests = '\n'.join(test_lines).strip() if test_lines else None317 318 # create bug summary from explanation319 detected = []320 if explanation:321 cleaned = explanation.replace("Explanation:", "").replace("\n", " ").strip()322 sentences = cleaned.split('.')323 if sentences and len(sentences[0].strip()) > 20:324 detected = [sentences[0].strip() + "."]325 326 if not detected:327 detected = ["Potential bug detected - see explanation for details"]328 329 return CodeResponse(330 text=raw,331 detectedBugs=detected,332 fixedCode=fixed_code,333 explanation=explanation,334 whyFix=None, # Your model doesn't separate this335 tests=tests,336 )337 338 # For chat or unknown tasks339 return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)340 341@app.get("/")342def read_root():343 return {344 "name": "Python Bug Fixer API",345 "status": "running",346 "models": {347 "base": "microsoft/phi-2",348 "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"349 },350 "version": "2.1.0",351 "info": "Shared base model approach - efficient memory usage"352 }353 354@app.get("/health")355def health_check():356 return {357 "status": "healthy",358 "base_model_loaded": base_model is not None,359 "finetuned_model_loaded": finetuned_model is not None,360 "device": str(device)361 }362 363 364@app.post("/analyze", response_model=CodeResponse)365async def analyze_code(request: CodeRequest): 366 try:367 print(f"Request: task={request.task}, use_base_model={request.use_base_model}")368 response_text = generate_response(369 request.code, 370 request.task, 371 request.use_base_model,372 request.temperature373 )374 return parse_response(response_text, request.task)375 except Exception as e:376 print(f"Error: {str(e)}")377 raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")378 379if __name__ == "__main__":380 uvicorn.run(app, host="0.0.0.0", port=7860)