grammar-inference-engine/bex/gbnf.py
tobjend 3468813ec8 fix: validate SORE before returning, skip malformed grammars
- Add validate_sore() to gbnf.py — checks parseability without converting
- _infer_group now validates grammar and returns skip_reason='malformed_grammar'
  for SOREs containing raw code (e.g. w_body=, sult=, (+,+:N+...)
- Results: 130 grammars, 130 GBNF OK, 0 GBNF FAIL
  - Flask: 5 OK, 0 FAIL, 3.4s
  - RAGSAK: 19 OK, 0 FAIL, 11 malformed, 12.9s
  - FastAPI: 106 OK, 0 FAIL, 6 malformed, 30.5s
2026-07-12 02:56:13 +02:00

385 lines
12 KiB
Python

"""GBNF — Convert SOREs to GBNF grammar format for llama.cpp constrained decoding.
GBNF (GGML BNF) is the de-facto standard grammar format for grammar-constrained
LLM output. It supports: literals, concatenation, alternation, repetition (+, *, ?),
and grouping.
SORE operators map directly:
SORE `.` (concat) → GBNF implicit concat (space)
SORE `|` (alt) → GBNF `|`
SORE `+` (plus) → GBNF `+`
SORE `?` (optional) → GBNF `?`
SORE `*` (star) → GBNF `*`
SORE `()` (group) → GBNF `()`
SORE literal → GBNF `"literal"`
"""
# ---------------------------------------------------------------------------
# SORE tokenizer
# ---------------------------------------------------------------------------
_TOKEN_CHARS = set('.|+?*()')
def _tokenize(sore):
"""Tokenize a SORE string into a list of (type, value) tuples.
Literals are sequences of non-operator characters.
"""
tokens = []
i = 0
n = len(sore)
while i < n:
ch = sore[i]
if ch == '.':
tokens.append(('DOT', ch))
i += 1
elif ch == '|':
tokens.append(('PIPE', ch))
i += 1
elif ch == '+':
tokens.append(('PLUS', ch))
i += 1
elif ch == '?':
tokens.append(('QUESTION', ch))
i += 1
elif ch == '*':
tokens.append(('STAR', ch))
i += 1
elif ch == '(':
tokens.append(('LPAREN', ch))
i += 1
elif ch == ')':
tokens.append(('RPAREN', ch))
i += 1
elif ch == 'ε':
tokens.append(('EPSILON', 'ε'))
i += 1
elif ch == '':
tokens.append(('EMPTY', ''))
i += 1
else:
# Collect literal characters (until next operator or paren)
start = i
while i < n and sore[i] not in _TOKEN_CHARS and sore[i] not in 'ε∅':
i += 1
lit = sore[start:i]
if lit:
# Strip newlines/extra whitespace from symbol names
lit = ' '.join(lit.split())
tokens.append(('LITERAL', lit))
return tokens
# ---------------------------------------------------------------------------
# SORE parser — produces a tree
# ---------------------------------------------------------------------------
class _Node:
"""AST node for SORE."""
pass
class _Literal(_Node):
def __init__(self, value):
self.value = value
def __repr__(self):
return f'Lit({self.value!r})'
class _Concat(_Node):
def __init__(self, parts):
self.parts = parts
def __repr__(self):
return f'Concat({self.parts})'
class _Alt(_Node):
def __init__(self, parts):
self.parts = parts
def __repr__(self):
return f'Alt({self.parts})'
class _Plus(_Node):
def __init__(self, child):
self.child = child
def __repr__(self):
return f'Plus({self.child})'
class _Optional(_Node):
def __init__(self, child):
self.child = child
def __repr__(self):
return f'Optional({self.child})'
class _Star(_Node):
def __init__(self, child):
self.child = child
def __repr__(self):
return f'Star({self.child})'
class _Epsilon(_Node):
def __repr__(self):
return 'Epsilon()'
class _Empty(_Node):
def __repr__(self):
return 'Empty()'
class _Parser:
"""Recursive descent parser for SOREs.
Handles the overloaded + operator:
- (a+b+c) → disjunction (inside parens)
- r+ → one-or-more repetition (outside parens)
"""
def __init__(self, tokens):
self.tokens = tokens
self.pos = 0
self.paren_depth = 0
def peek(self):
if self.pos < len(self.tokens):
return self.tokens[self.pos]
return ('EOF', '')
def consume(self, expected_type=None):
tok = self.peek()
if tok[0] == 'EOF':
raise ValueError(f'Unexpected end of SORE, expected {expected_type}')
if expected_type and tok[0] != expected_type:
raise ValueError(f'Expected {expected_type}, got {tok}')
self.pos += 1
return tok
def parse(self):
"""Parse full SORE."""
node = self.parse_alternation()
return node
def parse_alternation(self):
"""Parse: concat (('|' | '+') concat)* — + is alternation inside parens"""
parts = [self.parse_concat()]
while self.peek()[0] in ('PIPE', 'PLUS'):
if self.peek()[0] == 'PLUS' and self.paren_depth == 0:
break # + outside parens is repetition, not alternation
self.consume()
parts.append(self.parse_concat())
if len(parts) == 1:
return parts[0]
return _Alt(parts)
def parse_concat(self):
"""Parse: repetition (('.' | LPAREN) repetition)* — implicit concat"""
parts = [self.parse_repetition()]
while self.peek()[0] in ('DOT', 'LPAREN'):
if self.peek()[0] == 'DOT':
self.consume('DOT')
# LPAREN = implicit concat (no separator)
parts.append(self.parse_repetition())
if len(parts) == 1:
return parts[0]
return _Concat(parts)
def parse_repetition(self):
"""Parse: atom ('+' | '?' | '*')?
Inside parens, + is alternation (consumed by parse_alternation),
not repetition. Outside parens, always consume + as repetition.
Handles compound: +?, +*, ?+, *+ etc.
"""
node = self.parse_atom()
if self.peek()[0] in ('PLUS', 'QUESTION', 'STAR'):
if self.peek()[0] == 'PLUS' and self.paren_depth > 0:
return node # + inside parens is alternation, handled by caller
op = self.consume()
if op[0] == 'PLUS':
node = _Plus(node)
elif op[0] == 'QUESTION':
node = _Optional(node)
elif op[0] == 'STAR':
node = _Star(node)
# Handle compound repetition: +?, +*, ?+ etc.
# Normalize: Optional(Plus(x)) → Star(x)
if self.peek()[0] in ('PLUS', 'QUESTION', 'STAR'):
if self.peek()[0] == 'PLUS' and self.paren_depth > 0:
return node
op2 = self.consume()
if op2[0] == 'QUESTION':
if isinstance(node, _Plus):
node = _Star(node.child)
else:
node = _Optional(node)
elif op2[0] == 'STAR':
node = _Star(node.child if isinstance(node, (_Plus, _Optional)) else node)
elif op2[0] == 'PLUS':
if isinstance(node, _Optional):
node = _Plus(node.child)
elif isinstance(node, (_Plus, _Star)):
node = node # ++ is idempotent
else:
node = _Plus(node)
return node
def parse_atom(self):
"""Parse: literal | '(' alternation ')' | 'ε' | ''"""
tok = self.peek()
if tok[0] == 'EOF':
return _Epsilon()
if tok[0] == 'LITERAL':
self.consume()
return _Literal(tok[1])
if tok[0] == 'EPSILON':
self.consume()
return _Epsilon()
if tok[0] == 'EMPTY':
self.consume()
return _Empty()
if tok[0] == 'LPAREN':
self.consume('LPAREN')
self.paren_depth += 1
node = self.parse_alternation()
self.consume('RPAREN')
self.paren_depth -= 1
return node
raise ValueError(f'Unexpected token: {tok}')
def _parse_sore(sore):
"""Parse a SORE string into an AST."""
tokens = _tokenize(sore)
parser = _Parser(tokens)
return parser.parse()
# ---------------------------------------------------------------------------
# AST → GBNF
# ---------------------------------------------------------------------------
def _needs_group(node):
"""Check if a node needs parentheses in GBNF output."""
return isinstance(node, (_Alt, _Concat))
def _node_to_gbnf(node, rule_counter):
"""Convert AST node to GBNF fragment string.
Returns (gbnf_string, new_rule_counter, new_rules_list).
new_rules_list contains any helper rules needed.
"""
if isinstance(node, _Literal):
# Escape special chars in the literal for GBNF
escaped = node.value.replace('\\', '\\\\').replace('"', '\\"')
return f'"{escaped}"', rule_counter, []
if isinstance(node, _Epsilon):
return '', rule_counter, []
if isinstance(node, _Empty):
return '', rule_counter, []
if isinstance(node, _Concat):
parts = []
all_new_rules = []
for child in node.parts:
frag, rule_counter, new_rules = _node_to_gbnf(child, rule_counter)
if isinstance(child, _Alt):
frag = f'({frag})'
parts.append(frag)
all_new_rules.extend(new_rules)
return ' '.join(p for p in parts if p), rule_counter, all_new_rules
if isinstance(node, _Alt):
parts = []
all_new_rules = []
for child in node.parts:
frag, rule_counter, new_rules = _node_to_gbnf(child, rule_counter)
parts.append(frag)
all_new_rules.extend(new_rules)
return ' | '.join(p for p in parts if p), rule_counter, all_new_rules
if isinstance(node, _Plus):
frag, rule_counter, new_rules = _node_to_gbnf(node.child, rule_counter)
if _needs_group(node.child):
return f'({frag})+', rule_counter, new_rules
return f'{frag}+', rule_counter, new_rules
if isinstance(node, _Optional):
frag, rule_counter, new_rules = _node_to_gbnf(node.child, rule_counter)
if _needs_group(node.child):
return f'({frag})?', rule_counter, new_rules
return f'{frag}?', rule_counter, new_rules
if isinstance(node, _Star):
frag, rule_counter, new_rules = _node_to_gbnf(node.child, rule_counter)
if _needs_group(node.child):
return f'({frag})*', rule_counter, new_rules
return f'{frag}*', rule_counter, new_rules
raise ValueError(f'Unknown node type: {type(node)}')
def to_gbnf(sore):
"""Convert a SORE string to a GBNF grammar string.
Args:
sore: A SORE string like 'raise.(ValueError)+' or 'mockk'
Returns:
A GBNF grammar string with a single 'root' rule.
Example:
>>> to_gbnf('mockk')
'root ::= "mockk"'
>>> to_gbnf('raise.(ValueError)+')
'root ::= "raise" ("ValueError")+'
>>> to_gbnf('assertEquals.(of.(assertFailsWith)?)+')
'root ::= "assertEquals" ("of" ("assertFailsWith")?)+'
"""
if not sore or sore == '':
return 'root ::= ""'
if sore == 'ε':
return 'root ::= ""'
tree = _parse_sore(sore)
frag, _, _ = _node_to_gbnf(tree, 0)
return f'root ::= {frag}'
def to_gbnf_with_rules(sore, name='root'):
"""Convert a SORE to GBNF with a named rule.
Args:
sore: A SORE string
name: Rule name (default: 'root')
Returns:
GBNF rule string like 'my-rule ::= "foo" ("bar")+'
"""
if not sore or sore == '':
return f'{name} ::= ""'
if sore == 'ε':
return f'{name} ::= ""'
tree = _parse_sore(sore)
frag, _, _ = _node_to_gbnf(tree, 0)
return f'{name} ::= {frag}'
def validate_sore(sore):
"""Check if a SORE string is parseable. Returns (True, None) or (False, error_msg)."""
if not sore or sore in ('', 'ε'):
return True, None
try:
_parse_sore(sore)
return True, None
except ValueError as e:
return False, str(e)