File size: 5,684 Bytes
6a7acce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
import ast
import re
import numpy as np

MAX_INPUT_TEXT_LEN = 1000
MAX_TOTAL_ELEMENTS = 50000
MAX_ARRAY_RANK = 2
MAX_NUMERIC_MAGNITUDE = 1e9

def split_top_level_assignments(text):
    """
    Split an assignment string (e.g. 'W = [[1, 2], [3, 4]], x = [1, 1], b = 0.5')
    only at commas that are at bracket depth 0.
    """
    chunks = []
    current = []
    depth = 0
    for char in text:
        if char in '([{':
            depth += 1
            current.append(char)
        elif char in ')]}':
            depth -= 1
            if depth < 0:
                raise ValueError("Unmatched closing bracket in variable values.")
            current.append(char)
        elif char == ',' and depth == 0:
            chunks.append(''.join(current).strip())
            current = []
        else:
            current.append(char)
    if depth != 0:
        raise ValueError("Unmatched opening bracket in variable values.")
    last = ''.join(current).strip()
    if last:
        chunks.append(last)
    return chunks


def _validate_and_convert_literal(node):
    """
    Recursively inspect an AST literal node and convert to a Python float or nested list of floats.
    Strictly forbids function calls, comprehensions, variables, strings, booleans, etc.
    """
    if isinstance(node, ast.Constant):
        # In Python 3.8+, booleans are an instance of int/Constant with value bool
        if isinstance(node.value, bool):
            raise ValueError("Boolean literals are not allowed.")
        if not isinstance(node.value, (int, float)):
            raise ValueError(f"Only numeric literals are allowed, got {type(node.value).__name__}.")
        val = float(node.value)
        if np.isnan(val) or np.isinf(val):
            raise ValueError("Values must be finite real numbers.")
        if abs(val) > MAX_NUMERIC_MAGNITUDE:
            raise ValueError(f"Numeric literal exceeds maximum allowed magnitude ({MAX_NUMERIC_MAGNITUDE}).")
        return val

    if hasattr(ast, "Num") and isinstance(node, ast.Num):
        val = float(node.n)
        if abs(val) > MAX_NUMERIC_MAGNITUDE:
            raise ValueError(f"Numeric literal exceeds maximum allowed magnitude.")
        return val

    if isinstance(node, ast.UnaryOp):
        if isinstance(node.op, ast.UAdd):
            return +_validate_and_convert_literal(node.operand)
        elif isinstance(node.op, ast.USub):
            return -_validate_and_convert_literal(node.operand)
        raise ValueError(f"Unsupported unary operator in literal: {type(node.op).__name__}")

    if isinstance(node, ast.List):
        if len(node.elts) == 0:
            raise ValueError("Empty lists or arrays are not allowed.")
        return [_validate_and_convert_literal(elt) for elt in node.elts]

    raise ValueError(f"Disallowed construct in literal: {type(node).__name__}")


def parse_literal_value(val_str):
    """
    Safely parse a literal string (number, vector list, or matrix list) into a float or np.ndarray.
    Enforces rectangular matrix shape, finite elements, rank <= 2, and element count budget.
    """
    val_str = val_str.strip()
    if not val_str:
        raise ValueError("Empty value literal.")

    try:
        parsed = ast.parse(val_str, mode="eval")
    except SyntaxError as e:
        raise ValueError(f"Syntax error in literal '{val_str}': {e.msg}")

    val = _validate_and_convert_literal(parsed.body)

    if isinstance(val, (int, float)):
        return float(val)

    # Convert nested lists to numpy array
    try:
        arr = np.array(val, dtype=np.float64)
    except ValueError as e:
        raise ValueError(f"Malformed or ragged array structure: {e}")

    if arr.ndim > MAX_ARRAY_RANK:
        raise ValueError(f"Array rank {arr.ndim} exceeds maximum supported rank ({MAX_ARRAY_RANK}).")

    if arr.size == 0:
        raise ValueError("Empty array dimensions are not allowed.")

    if arr.size > MAX_TOTAL_ELEMENTS:
        raise ValueError(f"Array size ({arr.size} elements) exceeds limit ({MAX_TOTAL_ELEMENTS}).")

    if not np.all(np.isfinite(arr)):
        raise ValueError("Array contains nonfinite numbers (NaN or Inf).")

    return arr


def parse_variable_assignments(text):
    """
    Parse a comma-separated string of variable assignments into a dictionary:
        { "W": np.ndarray, "x": np.ndarray, "b": float, ... }
    Validates variable names, duplicate keys, and bounds.
    """
    if text is None or not text.strip():
        return {}

    cleaned = text.strip()
    if len(cleaned) > MAX_INPUT_TEXT_LEN:
        raise ValueError(f"Input text exceeds maximum length ({MAX_INPUT_TEXT_LEN} characters).")

    chunks = split_top_level_assignments(cleaned)
    feed_dict = {}
    total_elements = 0

    for chunk in chunks:
        if not chunk:
            continue
        if "=" not in chunk:
            raise ValueError(f"Expected assignment 'var = value', got '{chunk}'")

        var_part, val_part = chunk.split("=", 1)
        var_name = var_part.strip()
        val_str = val_part.strip()

        if not re.match(r"^[a-zA-Z][a-zA-Z0-9_]*$", var_name):
            raise ValueError(f"Invalid variable identifier: '{var_name}'")

        if var_name in feed_dict:
            raise ValueError(f"Duplicate assignment for variable '{var_name}'")

        parsed_val = parse_literal_value(val_str)
        if isinstance(parsed_val, np.ndarray):
            total_elements += parsed_val.size
        else:
            total_elements += 1

        if total_elements > MAX_TOTAL_ELEMENTS:
            raise ValueError(f"Total element count exceeds limit ({MAX_TOTAL_ELEMENTS}).")

        feed_dict[var_name] = parsed_val

    return feed_dict