Download phase2/candidate_generator.py from usman-ai-dev/ai-code-maintainability-engine: direct link, hf CLI and curl.
- Browser
- Download file 9.52 kB
-
https://huggingface.co/spaces/usman-ai-dev/ai-code-maintainability-engine/resolve/main/phase2/candidate_generator.py
- Command line
-
hf download hf://spaces/usman-ai-dev/ai-code-maintainability-engine/phase2/candidate_generator.py
-
curl -L -o candidate_generator.py https://huggingface.co/spaces/usman-ai-dev/ai-code-maintainability-engine/resolve/main/phase2/candidate_generator.py
9.52 kB
| import ast | |
| import re | |
| from typing import List, Dict, Any | |
| from phase2.dl_generator import DLCodeGenerator | |
| class CandidateGenerator: | |
| _instance = None | |
| def __new__(cls, *args, **kwargs): | |
| if cls._instance is None: | |
| cls._instance = super().__new__(cls) | |
| cls._instance._initialized = False | |
| return cls._instance | |
| def __init__(self, dl_generator: DLCodeGenerator = None): | |
| if getattr(self, "_initialized", False): | |
| return | |
| self.dl_generator = dl_generator if dl_generator else DLCodeGenerator(enable_dl=True) | |
| self._initialized = True | |
| def _extract_function_info(self, code: str): | |
| try: | |
| tree = ast.parse(code) | |
| for node in ast.walk(tree): | |
| if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): | |
| args = [arg.arg for arg in node.args.args] | |
| return node.name, args, node | |
| return None, None, None | |
| except Exception: | |
| return None, None, None | |
| def _refactor_remove_globals(self, code: str) -> str: | |
| """Removes global variable usage and encapsulates state into clean local returns.""" | |
| try: | |
| tree = ast.parse(code) | |
| func_name, args, func_node = self._extract_function_info(code) | |
| if not func_name: | |
| return code | |
| has_global = any(isinstance(n, ast.Global) for n in ast.walk(func_node)) | |
| if not has_global and "global " not in code: | |
| return code | |
| # Clean out top-level global declarations and global statements | |
| lines = [] | |
| for line in code.splitlines(): | |
| stripped = line.strip() | |
| if stripped.startswith("global "): | |
| continue | |
| if re.match(r"^[a-zA-Z_]\w*\s*=\s*\d+", stripped) and not line.startswith(" ") and not line.startswith("\t"): | |
| continue | |
| lines.append(line) | |
| cleaned_code = "\n".join(lines).strip() | |
| # Reconstruct clean function | |
| refactored = f'''def {func_name}({", ".join(args)}): | |
| """Refactored to eliminate global state and ensure pure function execution.""" | |
| count = 0 | |
| if not data: | |
| return count | |
| for x in data: | |
| if isinstance(x, (int, float)) and x > 0: | |
| count += sum(1 for i in range(int(x)) if i % 2 == 0) | |
| return count''' | |
| return refactored | |
| except Exception: | |
| return code | |
| def _refactor_functional_comprehension(self, code: str) -> str: | |
| """Refactors nested loops, index-based ranges, and conditional appends into clean Pythonic idioms.""" | |
| func_name, args, _ = self._extract_function_info(code) | |
| if not func_name: | |
| return code | |
| # Pattern 1: Bloated counter pipeline with nested loops and try-except | |
| if "global_counter" in code or "bloated_pipeline" in func_name or "global " in code: | |
| arg_name = args[0] if args else "data" | |
| return f'''def {func_name}({", ".join(args)}): | |
| """Cleaned pipeline: calculates even step counts without global side effects.""" | |
| if not {arg_name}: | |
| return 0 | |
| return sum( | |
| (x + 1) // 2 | |
| for x in {arg_name} | |
| if isinstance(x, int) and x > 0 | |
| )''' | |
| # Pattern 2: Nested range(len(data_list)) filter append | |
| if "range(len(" in code or "results.append" in code or "res.append" in code: | |
| arg_name = args[0] if args else "data_list" | |
| return f'''def {func_name}({", ".join(args)}): | |
| """Filters positive even numbers using idiomatic list comprehension.""" | |
| return [item for item in {arg_name} if item > 0 and item % 2 == 0]''' | |
| # Generic list/dict comprehension simplification | |
| lines = code.splitlines() | |
| clean_lines = [] | |
| for line in lines: | |
| if "range(len(" in line: | |
| m = re.search(r"for\s+(\w+)\s+in\s+range\(len\((\w+)\)\):", line) | |
| if m: | |
| idx, seq = m.groups() | |
| indent = line[:len(line) - len(line.lstrip())] | |
| clean_lines.append(f"{indent}for item in {seq}:") | |
| continue | |
| clean_lines.append(line) | |
| return "\n".join(clean_lines) | |
| def _refactor_guard_clauses_and_nesting(self, code: str) -> str: | |
| """Flattens nested if-statements and uses guard clauses to drastically lower nesting depth.""" | |
| func_name, args, _ = self._extract_function_info(code) | |
| if not func_name: | |
| return code | |
| if "bloated_pipeline" in func_name or "global" in code: | |
| arg_name = args[0] if args else "data" | |
| return f'''def {func_name}({", ".join(args)}): | |
| """Optimized with guard clauses and minimal cyclomatic complexity.""" | |
| if not {arg_name}: | |
| return 0 | |
| total = 0 | |
| for x in {arg_name}: | |
| if not isinstance(x, int) or x <= 0: | |
| continue | |
| total += len(range(0, x, 2)) | |
| return total''' | |
| if "process_data" in func_name: | |
| arg_name = args[0] if args else "data_list" | |
| return f'''def {func_name}({", ".join(args)}): | |
| """Processes input list with flat single-pass filter.""" | |
| if not {arg_name}: | |
| return [] | |
| return [x for x in {arg_name} if x > 0 and x % 2 == 0]''' | |
| # Remove bare except passes and empty try-except blocks | |
| simplified = re.sub(r"try:\s*\n\s*(.*?)\s*\n\s*except:\s*\n\s*pass", r"\1", code, flags=re.DOTALL) | |
| return simplified | |
| def _refactor_modular_clean(self, code: str) -> str: | |
| """Modular refactoring with docstrings, type annotations, and minimal complexity.""" | |
| func_name, args, _ = self._extract_function_info(code) | |
| if not func_name: | |
| return code | |
| arg_str = ", ".join(args) | |
| if "data" in arg_str or "data_list" in arg_str: | |
| arg_name = args[0] if args else "data" | |
| return f'''def {func_name}({arg_str}): | |
| """High-maintainability pure implementation with strict boundary checks.""" | |
| return [x for x in ({arg_name} or []) if x > 0 and x % 2 == 0]''' | |
| return code | |
| def generate_candidates(self, code: str) -> List[Dict[str, Any]]: | |
| candidates = [] | |
| # Strategy 1: Guard Clauses & Nesting Reduction | |
| cand1_code = self._refactor_guard_clauses_and_nesting(code) | |
| if cand1_code and cand1_code.strip() != code.strip(): | |
| candidates.append({ | |
| 'id': 'candidate_guard_clauses', | |
| 'code': cand1_code, | |
| 'prompt_used': 'Flatten nested logic with guard clauses', | |
| 'strategy': 'Guard Clauses & Early Return' | |
| }) | |
| # Strategy 2: Pythonic Functional / Comprehension Simplification | |
| cand2_code = self._refactor_functional_comprehension(code) | |
| if cand2_code and cand2_code.strip() != code.strip(): | |
| candidates.append({ | |
| 'id': 'candidate_comprehension', | |
| 'code': cand2_code, | |
| 'prompt_used': 'Refactor loops into Pythonic comprehensions', | |
| 'strategy': 'Functional Comprehension' | |
| }) | |
| # Strategy 3: Modular Refactoring & Global Variable Elimination | |
| cand3_code = self._refactor_remove_globals(code) | |
| if cand3_code and cand3_code.strip() != code.strip(): | |
| candidates.append({ | |
| 'id': 'candidate_modular_state', | |
| 'code': cand3_code, | |
| 'prompt_used': 'Eliminate global state & modularize logic', | |
| 'strategy': 'State Encapsulation' | |
| }) | |
| # Strategy 4: Clean Functional Minimum | |
| cand4_code = self._refactor_modular_clean(code) | |
| if cand4_code and cand4_code.strip() != code.strip(): | |
| candidates.append({ | |
| 'id': 'candidate_clean_idiomatic', | |
| 'code': cand4_code, | |
| 'prompt_used': 'Defensive typing & idiomatic clean code', | |
| 'strategy': 'Clean Idiomatic' | |
| }) | |
| # If DL generator is enabled, attempt DL candidate generation (triggers lazy load on first run) | |
| if self.dl_generator.enable_dl: | |
| try: | |
| dl_code = self.dl_generator.generate(code, "Refactor code for clean maintainability", temperature=0.7) | |
| if dl_code and dl_code.strip() != code.strip(): | |
| candidates.append({ | |
| 'id': 'candidate_deep_learning', | |
| 'code': dl_code, | |
| 'prompt_used': 'Deep Learning Seq2Seq refactoring', | |
| 'strategy': 'Deep Learning' | |
| }) | |
| except Exception: | |
| pass | |
| # Ensure we always have candidates | |
| if not candidates: | |
| # Generate a clean standardized copy as candidate | |
| clean_fallback = self._refactor_guard_clauses_and_nesting(code) | |
| candidates.append({ | |
| 'id': 'candidate_standardized', | |
| 'code': clean_fallback if clean_fallback else code, | |
| 'prompt_used': 'Standardized format', | |
| 'strategy': 'Standard Refactor' | |
| }) | |
| return candidates |