"""GBNF — Convert AST to GBNF grammar format for llama.cpp constrained decoding.""" from .grammar import ( Symbol, Concat, Alt, Plus, Optional, Star, Epsilon, Empty, ) def _needs_group(node): """Check if a node needs parentheses in GBNF output.""" return isinstance(node, (Alt, Concat)) def _node_to_gbnf(node): """Convert AST node to GBNF fragment string.""" if isinstance(node, Symbol): escaped = node.value.replace('\\', '\\\\').replace('"', '\\"') return f'"{escaped}"' if isinstance(node, (Epsilon, Empty)): return '' if isinstance(node, Concat): parts = [] for child in node.parts: frag = _node_to_gbnf(child) if isinstance(child, Alt): frag = f'({frag})' parts.append(frag) return ' '.join(p for p in parts if p) if isinstance(node, Alt): parts = [_node_to_gbnf(child) for child in node.parts] return ' | '.join(p for p in parts if p) if isinstance(node, Plus): frag = _node_to_gbnf(node.child) if _needs_group(node.child): return f'({frag})+' return f'{frag}+' if isinstance(node, Optional): frag = _node_to_gbnf(node.child) if _needs_group(node.child): return f'({frag})?' return f'{frag}?' if isinstance(node, Star): frag = _node_to_gbnf(node.child) if _needs_group(node.child): return f'({frag})*' return f'{frag}*' return '' def to_gbnf(node): """Convert AST node to a GBNF grammar string. Args: node: Grammar AST node Returns: GBNF grammar string with a single 'root' rule. """ if node is None or isinstance(node, Empty): return 'root ::= ""' if isinstance(node, Epsilon): return 'root ::= ""' frag = _node_to_gbnf(node) return f'root ::= {frag}' def to_gbnf_with_rules(node, name='root'): """Convert AST to GBNF with a named rule.""" if node is None or isinstance(node, Empty): return f'{name} ::= ""' if isinstance(node, Epsilon): return f'{name} ::= ""' frag = _node_to_gbnf(node) return f'{name} ::= {frag}' def grammar_structure_score(node): """Quantify how structured an AST is (0.0 = flat bag, 1.0 = fully structured).""" if node is None or isinstance(node, Empty): return 0.0 if isinstance(node, Symbol): return 0.0 if isinstance(node, Epsilon): return 0.0 if isinstance(node, (Plus, Optional, Star)): return 0.5 + 0.5 * grammar_structure_score(node.child) if isinstance(node, Alt): child_scores = [grammar_structure_score(p) for p in node.parts] return sum(child_scores) / max(len(child_scores), 1) if isinstance(node, Concat): n = len(node.parts) if n <= 1: return 1.0 child_scores = [grammar_structure_score(p) for p in node.parts] return min(1.0, 0.3 + 0.2 * n + sum(child_scores) / max(len(child_scores), 1)) return 0.0 # Noise token sets for grammar filtering TEST_NOISE = { 'assertEquals', 'assertTrue', 'assertFalse', 'assertNotNull', 'assertNull', 'every', 'verify', 'clearAllMocks', 'mockk', 'slot', 'coEvery', 'coVerify', 'assertThat', 'assertThrows', 'assertNotEquals', 'assumeTrue', 'doReturn', 'doThrow', 'assertSame', 'assertFailsWith', 'assertContains', 'runTest', 'TestRequest', 'TestClient', 'client', 'pytest', 'mock', 'patch', 'monkeypatch', 'tmp_path', 'async_client', 'test_client', } STDLIB_NOISE = { 'listOf', 'mapOf', 'setOf', 'arrayOf', 'mutableListOf', 'mutableMapOf', 'emptyList', 'emptyMap', 'emptySet', 'build', 'buildString', 'also', 'apply', 'let', 'run', 'to', 'of', 'get', 'set', 'if', 'else', 'when', 'return', 'is', 'in', 'as', 'toString', 'equals', 'hashCode', 'size', 'isEmpty', 'isNotEmpty', 'filter', 'map', 'flatMap', 'forEach', 'count', 'first', 'last', 'firstOrNull', 'single', 'singleOrNull', 'take', 'drop', 'joinToString', 'trim', 'isBlank', 'isNullOrBlank', 'orEmpty', 'contains', 'add', 'remove', 'clear', 'put', 'putAll', 'keys', 'values', 'String', 'Any', 'Boolean', 'Int', 'Long', 'Unit', 'Nothing', 'error', 'invoke', 'println', 'print', 'check', 'require', 'checkNotNull', 'requireNotNull', } # Combined noise set ALL_NOISE = TEST_NOISE | STDLIB_NOISE def filter_noise(node, noise_tokens=None): """Remove noise tokens from AST grammar. Walks the AST and removes Symbol nodes whose text is in the noise set. Returns cleaned AST, or Empty if everything was noise. Args: node: Grammar AST node noise_tokens: set of tokens to remove (default: ALL_NOISE) Returns: Cleaned AST node """ from .grammar import Concat, Alt, Optional, Plus, Star if noise_tokens is None: noise_tokens = ALL_NOISE if node is None or isinstance(node, (Epsilon, Empty)): return node if isinstance(node, Symbol): if node.value in noise_tokens: return Empty() return node if isinstance(node, Concat): new_parts = [] for part in node.parts: filtered = filter_noise(part, noise_tokens) if not isinstance(filtered, (Epsilon, Empty)): new_parts.append(filtered) if not new_parts: return Empty() if len(new_parts) == 1: return new_parts[0] return Concat(new_parts) if isinstance(node, Alt): new_parts = [] for part in node.parts: filtered = filter_noise(part, noise_tokens) if not isinstance(filtered, (Epsilon, Empty)): new_parts.append(filtered) if not new_parts: return Empty() if len(new_parts) == 1: return new_parts[0] return Alt(new_parts) if isinstance(node, (Plus, Optional, Star)): filtered = filter_noise(node.child, noise_tokens) if isinstance(filtered, (Epsilon, Empty)): return Empty() if isinstance(node, Plus): return Plus(filtered) if isinstance(node, Optional): return Optional(filtered) return Star(filtered) return node def grammar_noise_ratio(node, noise_tokens=None): """Calculate the fraction of symbols that are noise. Returns (n_noise, n_total) tuple. """ from .grammar import Concat, Alt, Optional, Plus, Star if noise_tokens is None: noise_tokens = ALL_NOISE if node is None or isinstance(node, (Epsilon, Empty)): return 0, 0 if isinstance(node, Symbol): is_noise = 1 if node.value in noise_tokens else 0 return is_noise, 1 if isinstance(node, Concat): noise = 0 total = 0 for part in node.parts: n, t = grammar_noise_ratio(part, noise_tokens) noise += n total += t return noise, total if isinstance(node, Alt): noise = 0 total = 0 for part in node.parts: n, t = grammar_noise_ratio(part, noise_tokens) noise += n total += t return noise, total if isinstance(node, (Plus, Optional, Star)): return grammar_noise_ratio(node.child, noise_tokens) return 0, 0