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