""" RegexFeatureExtractor - token feature gating. Extends build_code_features (4-dim) to a 7-dim per-token feature vector plus a running SyntaxState (bracket stack, quote open/closed) used as a cheap pre-model sanity gate. Dimensions per token: [0] is_code_like (indent / brackets / operators / newlines) [1] indent_depth (normalized leading whitespace) [2] bracket_balance (1.0 open, 0.5 neutral, 0.0 close) [3] has_newline [4] keyword_hit (NEW) [5] quote_state (running open/closed string literal) (NEW) [6] numeric_literal (NEW) """ import re import torch from dataclasses import dataclass, field from typing import List, Optional, Tuple CODE_CHARS = set("{}()[];=<>!&|+-*/%'\"`#@.,:") KEYWORD_RE = re.compile( r"\b(def|class|import|return|if|else|elif|for|while|try|except|finally|" r"lambda|pass|with|as|yield|from|async|await)\b" ) NUMERIC_RE = re.compile(r"\b\d+(\.\d+)?\b") OPEN_BRACKETS = {"{": 1, "[": 1, "(": 1} CLOSE_BRACKETS = {"}": 1, "]": 1, ")": 1} @dataclass class SyntaxState: bracket_stack: List[str] = field(default_factory=list) quote_open: Optional[str] = None # '"' or "'" while a string literal is open depth: int = 0 def is_balanced(self) -> bool: return not self.bracket_stack and self.quote_open is None def to_dict(self) -> dict: return { "balanced": self.is_balanced(), "bracket_depth": len(self.bracket_stack), "quote_open": self.quote_open, } class RegexFeatureExtractor: def __init__(self, num_features: int = 7): self.num_features = num_features def extract(self, tokenizer, input_ids: torch.Tensor) -> Tuple[torch.Tensor, List[SyntaxState]]: """Per-token features (B, T, F) + one SyntaxState per row.""" feats: List[List[List[float]]] = [] states: List[SyntaxState] = [] for row in input_ids.tolist(): tokens = tokenizer.convert_ids_to_tokens(row) state = SyntaxState() row_feats = [] for tok in tokens: is_code = any(c in CODE_CHARS for c in tok) indent = 0.0 stripped = tok.lstrip() if stripped and tok != stripped: indent = min((len(tok) - len(stripped)) / 8.0, 1.0) is_code = True bal = 0.5 for c in tok: if c in OPEN_BRACKETS: bal = 1.0 state.bracket_stack.append(c) elif c in CLOSE_BRACKETS: bal = 0.0 if state.bracket_stack: state.bracket_stack.pop() # running quote state for c in tok: if c in ('"', "'"): if state.quote_open is None: state.quote_open = c elif state.quote_open == c: state.quote_open = None quote = 1.0 if state.quote_open is not None else 0.0 newline = 1.0 if "\n" in tok else 0.0 kw = 1.0 if KEYWORD_RE.search(tok) else 0.0 num = 1.0 if NUMERIC_RE.search(tok) else 0.0 row_feats.append([1.0 if is_code else 0.0, indent, bal, newline, kw, quote, num]) state.depth = len(state.bracket_stack) row_feats = row_feats[: input_ids.shape[1]] feats.append(row_feats) states.append(state) max_len = max(len(r) for r in feats) padded = [ r + [[0.0, 0.0, 0.5, 0.0, 0.0, 0.0, 0.0]] * (max_len - len(r)) for r in feats ] t = torch.tensor(padded, dtype=torch.float32) if t.shape[-1] > self.num_features: t = t[..., : self.num_features] return t, states def extract_text(self, text: str) -> SyntaxState: """Run a string-only pass for the pre-model gate (no tokenizer).""" state = SyntaxState() for c in text: if c in OPEN_BRACKETS: state.bracket_stack.append(c) elif c in CLOSE_BRACKETS and state.bracket_stack: state.bracket_stack.pop() elif c in ('"', "'"): if state.quote_open is None: state.quote_open = c elif state.quote_open == c: state.quote_open = None state.depth = len(state.bracket_stack) return state def gate(self, syntax_state: SyntaxState) -> str: """Returns 'pass' | 'warn' | 'block'.""" if syntax_state.is_balanced(): return "pass" return "block" if syntax_state.depth > 4 else "warn" # drop-in replacement for architecture.build_code_features with 7-dim output def build_code_features_v2(tokenizer, input_ids: torch.Tensor) -> torch.Tensor: return RegexFeatureExtractor(num_features=7).extract(tokenizer, input_ids)[0]