File size: 2,047 Bytes
38bc0dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 | import ast
from typing import List, Dict
class CodeValidator:
def __init__(self, min_length_ratio: float = 0.15):
self.min_length_ratio = min_length_ratio
def _extract_function_signature(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)
return (None, None)
except Exception:
return (None, None)
def is_valid(self, original_code: str, candidate_code: str) -> bool:
if not candidate_code or not candidate_code.strip():
return False
# Must be valid Python syntax
try:
ast.parse(candidate_code)
except SyntaxError:
return False
# Check function signature compatibility if original has a function
orig_name, orig_args = self._extract_function_signature(original_code)
cand_name, cand_args = self._extract_function_signature(candidate_code)
if orig_name and cand_name:
if orig_name != cand_name:
return False
# Allow positional args match
if orig_args and cand_args:
if len(orig_args) != len(cand_args):
return False
orig_has_return = "return " in original_code or "return\n" in original_code
cand_has_return = "return " in candidate_code or "return\n" in candidate_code
if orig_has_return and not cand_has_return:
return False
return True
def filter_valid_candidates(self, original_code: str, candidates: List[Dict[str, str]]) -> List[Dict[str, str]]:
valid_candidates = []
for cand in candidates:
if self.is_valid(original_code, cand.get('code', '')):
valid_candidates.append(cand)
return valid_candidates
|