Spaces:
Running
Running
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
|