Research source · python

check_left_quotient.py

site/public/research-artifacts/left-quotient-regular-expressions/check_left_quotient.py

368 lines. Source is displayed for inspection; it is not executed by this page.

File fingerprint

SHA-256: 0d5773778c5c7e4e0b709385f4a2617fdf4aa1ae0179e83e0353da68c8163075

#!/usr/bin/env python3
"""Finite audit for the derivative-carrier left-quotient construction.

The construction under test interprets the quotienting expression A as a
relation on the normalized Brzozowski-derivative DFA of B.  The independent
oracle searches the synchronized derivative product of A and B for a prefix
accepted by A, then tests the supplied suffix in the reached B state.
"""

from __future__ import annotations

import hashlib
import itertools
import json
from collections import deque
from functools import lru_cache

ALPHABET = ("a", "b")
ZERO = ("zero",)
ONE = ("one",)


def lit(letter: str):
    return ("lit", letter)


def alt(*expressions):
    terms = []
    for expression in expressions:
        if expression == ZERO:
            continue
        if expression[0] == "alt":
            terms.extend(expression[1:])
        else:
            terms.append(expression)
    terms = sorted(set(terms), key=repr)
    if not terms:
        return ZERO
    if len(terms) == 1:
        return terms[0]
    return ("alt", *terms)


def cat(*expressions):
    factors = []
    for expression in expressions:
        if expression == ZERO:
            return ZERO
        if expression == ONE:
            continue
        if expression[0] == "cat":
            factors.extend(expression[1:])
        else:
            factors.append(expression)
    if not factors:
        return ONE
    if len(factors) == 1:
        return factors[0]
    return ("cat", *factors)


def star(expression):
    if expression in (ZERO, ONE):
        return ONE
    if expression[0] == "star":
        return expression
    return ("star", expression)


@lru_cache(maxsize=None)
def nullable(expression) -> bool:
    tag = expression[0]
    if tag == "zero" or tag == "lit":
        return False
    if tag == "one" or tag == "star":
        return True
    if tag == "alt":
        return any(nullable(term) for term in expression[1:])
    if tag == "cat":
        return all(nullable(factor) for factor in expression[1:])
    raise ValueError(tag)


@lru_cache(maxsize=None)
def derivative(expression, letter: str):
    tag = expression[0]
    if tag in ("zero", "one"):
        return ZERO
    if tag == "lit":
        return ONE if expression[1] == letter else ZERO
    if tag == "alt":
        return alt(*(derivative(term, letter) for term in expression[1:]))
    if tag == "cat":
        factors = expression[1:]
        terms = []
        prefix_nullable = True
        for index, factor in enumerate(factors):
            if prefix_nullable:
                terms.append(cat(derivative(factor, letter), *factors[index + 1 :]))
            prefix_nullable = prefix_nullable and nullable(factor)
        return alt(*terms)
    if tag == "star":
        body = expression[1]
        return cat(derivative(body, letter), star(body))
    raise ValueError(tag)


def run(expression, word: str):
    for letter in word:
        expression = derivative(expression, letter)
    return expression


def derivative_dfa(start):
    states = {start}
    edges = {}
    queue = deque([start])
    while queue:
        state = queue.popleft()
        for letter in ALPHABET:
            target = derivative(state, letter)
            edges[(state, letter)] = target
            if target not in states:
                states.add(target)
                queue.append(target)
    return states, edges


def compose(left, right):
    by_middle = {}
    for middle, target in right:
        by_middle.setdefault(middle, set()).add(target)
    return {
        (source, target)
        for source, middle in left
        for target in by_middle.get(middle, ())
    }


def reflexive_transitive_closure(relation, states):
    closure = set(relation) | {(state, state) for state in states}
    changed = True
    while changed:
        expanded = closure | compose(closure, closure)
        changed = expanded != closure
        closure = expanded
    return closure


def relation_semantics(expression, states, edges, mutation=None):
    identity = {(state, state) for state in states}

    @lru_cache(maxsize=None)
    def interpret(node):
        tag = node[0]
        if tag == "zero":
            return frozenset()
        if tag == "one":
            return frozenset(identity)
        if tag == "lit":
            letter = node[1]
            if mutation == "reverse_literal_edges":
                return frozenset((edges[(state, letter)], state) for state in states)
            return frozenset((state, edges[(state, letter)]) for state in states)
        if tag == "alt":
            return frozenset().union(*(interpret(term) for term in node[1:]))
        if tag == "cat":
            result = identity
            for factor in node[1:]:
                if mutation == "reverse_concatenation":
                    result = compose(interpret(factor), result)
                else:
                    result = compose(result, interpret(factor))
            return frozenset(result)
        if tag == "star":
            if mutation == "truncate_star_closure":
                return frozenset(identity | set(interpret(node[1])))
            return frozenset(
                reflexive_transitive_closure(interpret(node[1]), states)
            )
        raise ValueError(tag)

    return interpret(expression)


def quotient_construction_accepts(
    a_expression, b_expression, suffix: str, mutation=None
):
    b_states, b_edges = derivative_dfa(b_expression)
    relation = relation_semantics(a_expression, b_states, b_edges, mutation)
    if mutation == "forget_distinguished_start":
        residuals = {target for _, target in relation}
    else:
        residuals = {target for source, target in relation if source == b_expression}
    return any(nullable(run(residual, suffix)) for residual in residuals)


def product_oracle_accepts(a_expression, b_expression, suffix: str):
    reachable = {(a_expression, b_expression)}
    queue = deque(reachable)
    while queue:
        a_state, b_state = queue.popleft()
        for letter in ALPHABET:
            target = (derivative(a_state, letter), derivative(b_state, letter))
            if target not in reachable:
                reachable.add(target)
                queue.append(target)
    return any(
        nullable(a_state) and nullable(run(b_state, suffix))
        for a_state, b_state in reachable
    )


a = lit("a")
b = lit("b")

A_EXPRESSIONS = [
    ZERO,
    ONE,
    a,
    b,
    alt(a, b),
    cat(a, b),
    cat(b, a),
    star(a),
    star(b),
    star(alt(a, b)),
    cat(star(a), b),
    cat(a, star(b)),
    star(cat(a, b)),
    star(cat(b, a)),
    cat(alt(a, b), a),
    cat(a, alt(a, b)),
    alt(cat(a, b), cat(b, a)),
    cat(star(a), star(b)),
    star(cat(star(a), b)),
    cat(star(cat(a, b)), a),
    alt(ONE, cat(a, b)),
    cat(alt(ONE, a), b),
    star(alt(ONE, cat(a, b))),
    cat(star(alt(a, b)), cat(a, b)),
    alt(star(a), star(b)),
    cat(star(a), b, star(a)),
    star(cat(alt(a, b), alt(a, b))),
]

B_EXPRESSIONS = [
    ZERO,
    ONE,
    a,
    b,
    alt(a, b),
    cat(a, b),
    cat(b, a),
    star(a),
    star(b),
    star(alt(a, b)),
    cat(star(a), b),
    cat(a, star(b)),
    star(cat(a, b)),
    alt(ONE, cat(a, b)),
    cat(star(a), b, star(a)),
]


MUTATIONS = (
    "reverse_concatenation",
    "truncate_star_closure",
    "reverse_literal_edges",
    "forget_distinguished_start",
)


def words_up_to(maximum_length: int):
    return [
        "".join(letters)
        for length in range(maximum_length + 1)
        for letters in itertools.product(ALPHABET, repeat=length)
    ]


def main() -> int:
    assert len(A_EXPRESSIONS) == len(set(A_EXPRESSIONS)) == 27
    assert len(B_EXPRESSIONS) == len(set(B_EXPRESSIONS)) == 15

    suffixes = words_up_to(4)
    mismatches = []
    transcript = hashlib.sha256()
    mutation_transcript = hashlib.sha256()
    mutation_results = {
        mutation: {"mismatch_count": 0, "first_witness": None}
        for mutation in MUTATIONS
    }
    max_a_states = 0
    max_b_states = 0

    for a_index, a_expression in enumerate(A_EXPRESSIONS):
        a_states, _ = derivative_dfa(a_expression)
        max_a_states = max(max_a_states, len(a_states))
        for b_index, b_expression in enumerate(B_EXPRESSIONS):
            b_states, _ = derivative_dfa(b_expression)
            max_b_states = max(max_b_states, len(b_states))
            for suffix in suffixes:
                construction = quotient_construction_accepts(
                    a_expression, b_expression, suffix
                )
                oracle = product_oracle_accepts(a_expression, b_expression, suffix)
                transcript.update(
                    f"{a_index}:{b_index}:{suffix}:{int(construction)}:{int(oracle)}\n".encode()
                )
                if construction != oracle:
                    mismatches.append(
                        {
                            "a_index": a_index,
                            "b_index": b_index,
                            "suffix": suffix,
                            "construction": construction,
                            "oracle": oracle,
                        }
                    )
                for mutation in MUTATIONS:
                    mutated = quotient_construction_accepts(
                        a_expression, b_expression, suffix, mutation
                    )
                    mutation_transcript.update(
                        f"{mutation}:{a_index}:{b_index}:{suffix}:{int(mutated)}:{int(oracle)}\n".encode()
                    )
                    if mutated != oracle:
                        result = mutation_results[mutation]
                        result["mismatch_count"] += 1
                        if result["first_witness"] is None:
                            result["first_witness"] = {
                                "a_index": a_index,
                                "b_index": b_index,
                                "suffix": suffix,
                                "mutated_construction": mutated,
                                "oracle": oracle,
                            }

    all_mutations_detected = all(
        result["mismatch_count"] > 0 for result in mutation_results.values()
    )

    certificate = {
        "certificate_type": "fprd-left-quotient-finite-audit.v2",
        "alphabet": list(ALPHABET),
        "quotienting_expressions": len(A_EXPRESSIONS),
        "target_expressions": len(B_EXPRESSIONS),
        "expression_pairs": len(A_EXPRESSIONS) * len(B_EXPRESSIONS),
        "suffixes": len(suffixes),
        "maximum_suffix_length": 4,
        "membership_comparisons": len(A_EXPRESSIONS) * len(B_EXPRESSIONS) * len(suffixes),
        "mismatch_count": len(mismatches),
        "maximum_quotienting_dfa_states": max_a_states,
        "maximum_target_dfa_states": max_b_states,
        "transcript_sha256": transcript.hexdigest(),
        "mutation_controls": len(MUTATIONS),
        "all_mutations_detected": all_mutations_detected,
        "mutation_transcript_sha256": mutation_transcript.hexdigest(),
        "mutation_results": mutation_results,
        "mismatches": mismatches,
    }
    print(json.dumps(certificate, indent=2, sort_keys=True))
    return 1 if mismatches or not all_mutations_detected else 0


if __name__ == "__main__":
    raise SystemExit(main())