diff --git a/src/mcore_bridge/config/model_config.py b/src/mcore_bridge/config/model_config.py index 942f2cf..6234724 100644 --- a/src/mcore_bridge/config/model_config.py +++ b/src/mcore_bridge/config/model_config.py @@ -16,17 +16,102 @@ logger = get_logger() +# Layer-pattern strings are list literals with repetition, for example +# "([0]*3+[1]*1)*3". The alphabet is tiny, but eval still accepts "[0]*10**9" +# and will try to allocate it. Parse the grammar instead and refuse a result +# longer than a model would ever have. +_MAX_PATTERN_ITEMS = 100_000 + -# code borrowed from NVIDIA/Megatron-LM def _eval_pattern(pattern): - """ Validate and evaluate a string containing a Python list expression """ + """Parse a layer-pattern list expression. No eval.""" assert isinstance(pattern, str) - - # validate input, only allow comma, digits, [, ], (, ), +, and * - if bool(re.compile(r'[^,\d\[\]\(\)\+\*]').search(pattern)): + compact = re.sub(r'\s+', '', pattern) + if compact == '' or re.search(r'[^,\d\[\]\(\)\+\*]', compact): raise ValueError(f'Invalid pattern: {pattern}') - - return eval(pattern) + if '**' in compact: + raise ValueError(f'Invalid pattern: {pattern}') + parser = _PatternParser(compact) + value = parser.parse() + if parser.i != len(compact): + raise ValueError(f'Invalid pattern: {pattern}') + if not isinstance(value, list): + raise ValueError(f'Invalid pattern: {pattern}') + return value + + +class _PatternParser: + + def __init__(self, text): + self.text = text + self.i = 0 + + def parse(self): + return self._expr() + + def _peek(self): + return self.text[self.i] if self.i < len(self.text) else '' + + def _eat(self, char): + if self._peek() != char: + raise ValueError(f'Invalid pattern: {self.text}') + self.i += 1 + + def _expr(self): + value = self._term() + while self._peek() == '+': + self._eat('+') + right = self._term() + if not isinstance(value, list) or not isinstance(right, list): + raise ValueError(f'Invalid pattern: {self.text}') + value = value + right + if len(value) > _MAX_PATTERN_ITEMS: + raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}') + return value + + def _term(self): + value = self._atom() + if self._peek() == '*': + self._eat('*') + count = self._number() + if not isinstance(value, list): + raise ValueError(f'Invalid pattern: {self.text}') + if count > _MAX_PATTERN_ITEMS or len(value) * count > _MAX_PATTERN_ITEMS: + raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}') + value = value * count + return value + + def _atom(self): + if self._peek() == '(': + self._eat('(') + value = self._expr() + self._eat(')') + return value + if self._peek() == '[': + return self._list() + raise ValueError(f'Invalid pattern: {self.text}') + + def _list(self): + self._eat('[') + if self._peek() == ']': + self._eat(']') + return [] + items = [self._number()] + while self._peek() == ',': + self._eat(',') + items.append(self._number()) + if len(items) > _MAX_PATTERN_ITEMS: + raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}') + self._eat(']') + return items + + def _number(self): + start = self.i + if not self._peek().isdigit(): + raise ValueError(f'Invalid pattern: {self.text}') + while self._peek().isdigit(): + self.i += 1 + return int(self.text[start:self.i]) # code borrowed from NVIDIA/Megatron-LM diff --git a/tests/test_layer_pattern.py b/tests/test_layer_pattern.py new file mode 100644 index 0000000..ee68008 --- /dev/null +++ b/tests/test_layer_pattern.py @@ -0,0 +1,35 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Layer patterns are parsed, not eval'd, and huge repetitions are refused.""" +import ast +from pathlib import Path + + +def _parser(): + path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/config/model_config.py' + tree = ast.parse(path.read_text()) + body = [ + node for node in tree.body if (isinstance(node, ast.Assign) and any( + isinstance(target, ast.Name) and target.id == '_MAX_PATTERN_ITEMS' for target in node.targets)) or ( + isinstance(node, (ast.FunctionDef, ast.ClassDef)) and node.name in {'_eval_pattern', '_PatternParser'}) + ] + module = ast.Module(body=body, type_ignores=[]) + ast.fix_missing_locations(module) + namespace = {'re': __import__('re')} + exec(compile(module, str(path), 'exec'), namespace) + return namespace['_eval_pattern'] + + +def test_documented_pattern_matches_the_expanded_list(): + parse = _parser() + assert parse('([0]*3+[1]*1)*3') == [0, 0, 0, 1] * 3 + assert parse('([1]+[0]*2)') == [1, 0, 0] + + +def test_exponent_and_huge_repeat_are_rejected(): + parse = _parser() + for pattern in ('[0]*10**9', '[0]*(10**9)', '[0]*1000000'): + try: + parse(pattern) + except ValueError: + continue + raise AssertionError(pattern)