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