grammar-inference-engine/bex/soa.py

225 lines
6.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""SOA — Single Occurrence Automaton (Definition 6, TODS 2010).
Labels are grammar.py AST nodes (Symbol, Concat, Alt, Plus, etc.).
"""
import copy
from .grammar import Symbol, Concat, Alt, Plus, Optional, Star, Epsilon, Empty
class SOA:
"""
Node-labeled automaton (Definition 6, TODS 2010).
V = {src, sink} symbol-labeled states.
E ⊆ V × V, unlabeled edges.
Walk src=v₁,v₂,...,vₙ₊₁=sink accepts word lab(v₂)...lab(vₙ).
Labels are AST nodes from grammar.py.
"""
def __init__(self):
self._next = 0
self._succ = {}
self._pred = {}
self._label = {}
self.src = self._new()
self.sink = self._new()
def _new(self):
n = self._next
self._next += 1
self._succ[n] = set()
self._pred[n] = set()
self._label[n] = None
return n
def add_state(self, label):
n = self._new()
if isinstance(label, str):
label = Symbol(label)
self._label[n] = label
return n
def add_edge(self, f, t):
self._succ[f].add(t)
self._pred[t].add(f)
def rm_edge(self, f, t):
self._succ[f].discard(t)
self._pred[t].discard(f)
def rm_state(self, n):
if n in (self.src, self.sink):
return
for p in list(self._pred[n]):
self.rm_edge(p, n)
for s in list(self._succ[n]):
self.rm_edge(n, s)
del self._label[n]
del self._succ[n]
del self._pred[n]
def label(self, n):
return self._label.get(n)
def set_label(self, n, lab):
if isinstance(lab, str):
lab = Symbol(lab)
self._label[n] = lab
def succ(self, n):
return set(self._succ.get(n, set()))
def pred(self, n):
return set(self._pred.get(n, set()))
def has_edge(self, f, t):
return t in self._succ.get(f, set())
def states(self):
return [n for n in self._succ if n not in (self.src, self.sink) and self._label.get(n) is not None]
def count_symbol(self, sym):
"""Count states whose label base matches sym (string or Symbol).
Strips _N suffixes before comparing."""
import re
if isinstance(sym, Symbol):
target = sym.value
else:
target = sym
count = 0
for n, lab in self._label.items():
if n in (self.src, self.sink):
continue
if isinstance(lab, Symbol):
base = re.sub(r'_\d+$', '', lab.value)
if base == target:
count += 1
elif isinstance(lab, str):
base = re.sub(r'_\d+$', '', lab)
if base == target:
count += 1
return count
def _pred_plus(self, n):
r = set(self._pred.get(n, set()))
lab = self._label.get(n)
if isinstance(lab, (Plus, Star)):
r.add(n)
return r
def _succ_plus(self, n):
r = set(self._succ.get(n, set()))
lab = self._label.get(n)
if isinstance(lab, (Plus, Star)):
r.add(n)
return r
def copy(self):
return copy.deepcopy(self)
def accept(self, w):
cur = {self.src}
for sym in w:
nxt = set()
for s in cur:
for t in self._succ.get(s, set()):
lab = self._label.get(t)
if isinstance(lab, Symbol) and lab.value == sym:
nxt.add(t)
elif isinstance(lab, str) and lab == sym:
nxt.add(t)
if not nxt:
return False
cur = nxt
return any(self.sink in self._succ.get(s, set()) for s in cur)
def sink_reachable(self):
seen = set()
q = [self.src]
while q:
s = q.pop()
if s == self.sink:
return True
if s in seen:
continue
seen.add(s)
q.extend(self._succ.get(s, []))
return False
def num_non_special(self):
return sum(1 for n in self._succ if n not in (self.src, self.sink))
def is_final(self):
ns = self.states()
return len(ns) == 1 and self.has_edge(self.src, ns[0]) and self.has_edge(ns[0], self.sink)
def expression(self):
if not self.is_final():
return None
return self._label[self.states()[0]]
def contract(self, r, s, new_label):
"""
State contraction G[r,s ⇒ t] (Definition 11, TODS 2010).
"""
if isinstance(new_label, str):
new_label = Symbol(new_label)
t = self._new()
self._label[t] = new_label
for v in self._pred.get(r, set()) - {r, s}:
self.add_edge(v, t)
for v in self._pred.get(s, set()) - {r, s}:
self.add_edge(v, t)
for w in self._succ.get(r, set()) - {r, s}:
self.add_edge(t, w)
for w in self._succ.get(s, set()) - {r, s}:
self.add_edge(t, w)
if r in self._succ.get(s, set()):
self.add_edge(t, t)
self.rm_state(r)
self.rm_state(s)
return t
def contract_single(self, r, new_label):
"""Single-state substitution G[r ⇒ t] (Definition 11 note)."""
if isinstance(new_label, str):
new_label = Symbol(new_label)
if r in (self.src, self.sink):
return r
t = self._new()
self._label[t] = new_label
for v in self._pred.get(r, set()) - {r}:
self.add_edge(v, t)
for w in self._succ.get(r, set()) - {r}:
self.add_edge(t, w)
if r in self._succ.get(r, set()):
self.add_edge(t, t)
self.rm_state(r)
return t
def epsilon_closure(self):
"""G* (Definition 25, TODS 2010). Add self-loops for + states and ε-transitive closure."""
G = self.copy()
changed = True
while changed:
changed = False
for n in list(G._succ.keys()):
lab = G._label.get(n)
if isinstance(lab, (Plus, Star)):
if not G.has_edge(n, n):
G.add_edge(n, n)
changed = True
for n in list(G._succ.keys()):
for m in list(G._succ.get(n, set())):
mlab = G._label.get(m)
if isinstance(mlab, Epsilon):
for mp in list(G._succ.get(m, set())):
if mp != n and not G.has_edge(n, mp):
G.add_edge(n, mp)
changed = True
return G
def __repr__(self):
return f"SOA(nodes={len(self._succ)}, special={self.num_non_special()})"