#!/usr/bin/env python3
from collections import defaultdict
from typing import overload
from lark import Lark
from fosf.config import FOSF_GRAMMAR
from fosf.parsers.base import BaseOSFParser, _BaseOSFTransformer
from fosf.parsers.taxonomy import _TaxonomyTransformer
from fosf.syntax.base import Tag
from fosf.syntax.taxonomy import SortTaxonomy
from fosf.syntax.terms import DisjunctiveTerm, NormalTerm, Term
class _OsfTermTransformer(_BaseOSFTransformer):
def __init__(self):
super().__init__()
self.tags = set()
def subterm(self, tree):
"Process a subterm: FEATURE -> term."
feat = tree[0]
term = tree[1]
return (feat, term)
def subterms(self, tree):
"Process subterms: '(' subterm (',' subterm)* ')'."
d = defaultdict(list)
for k, v in tree:
d[k].append(v)
return d
def untagged_term(self, tree):
"""Process an untagged term: sort (subterms)?."""
sort = tree[0]
if len(tree) > 1:
return {"tag": None, "sort": sort, "subterms": tree[1]}
return {"tag": None, "sort": sort, "subterms": None}
def unsorted_term(self, tree):
"""Process an unsorted term: TAG (subterms)?."""
tag = Tag(tree[0].value)
self.tags.add(tag)
if len(tree) > 1:
return {"tag": tag, "sort": None, "subterms": tree[1]}
return {"tag": tag, "sort": None, "subterms": None}
def untagged_disjunctive_term(self, tree):
return {"tag": None, "disjunctive_terms": tree}
def disjunctive_term(self, tree):
tag = Tag(tree[0].value)
return {"tag": tag, "disjunctive_terms": tree[1:]}
def term(self, tree):
"""Process a tagged term: TAG (":" untagged_term)."""
tag = Tag(tree[0].value)
sort = tree[1]
self.tags.add(tag)
if len(tree) > 2:
return {"tag": tag, "sort": sort, "subterms": tree[2]}
return {"tag": tag, "sort": sort, "subterms": None}
def prefixed_term(self, tree):
return tree[-1]
def transform(self, parse_tree):
self.tags = set()
return super().transform(parse_tree)
[docs]
class OsfTermParser(BaseOSFParser):
def __init__(self):
self.parser = Lark.open_from_package(
"fosf", FOSF_GRAMMAR, start="prefixed_term"
)
self.transformer = _OsfTermTransformer()
self.term_constructor = Term
self.tags = set()
self.tag_counter = 0
def _dict_to_term(self, term_dict, default_tag):
def visit(term):
tag = self.__find_tag(default_tag) if term["tag"] is None else term["tag"]
if "disjunctive_terms" in term:
return DisjunctiveTerm(
tag, [visit(t) for t in term["disjunctive_terms"]]
)
sort = term["sort"]
if self.term_constructor == NormalTerm:
if term["subterms"] is None:
subterms = {}
else:
subterms = {k: visit(v[0]) for k, v in term["subterms"].items()}
else:
if term["subterms"] is None:
subterms = defaultdict(list)
else:
subterms = defaultdict(list)
for feature, values in term["subterms"].items():
for value in values:
subterms[feature].append(visit(value))
return self.term_constructor(tag, sort, subterms)
return visit(term_dict)
def __find_tag(self, default_tag):
while (tag := Tag(f"{default_tag}{self.tag_counter}")) in self.tags:
self.tag_counter += 1
self.tags.add(tag)
return tag
@overload
def parse(
self, expression: str, default_tag="X", create_using=NormalTerm
) -> NormalTerm: ...
@overload
def parse(self, expression: str, default_tag="X", create_using=Term) -> Term: ...
[docs]
def parse(self, expression: str, default_tag="X", create_using=None) -> Term:
parse_tree = self.parser.parse(expression)
if create_using is None:
self.term_constructor = Term
else:
self.term_constructor = create_using
term_dict = self.transformer.transform(parse_tree)
self.tags = self.transformer.tags
self.tag_counter = 0
return self._dict_to_term(term_dict, default_tag)
class _QueryTermTransformer(_OsfTermTransformer):
def q_tag(self, tree):
tag = Tag(tree[-1].value)
if len(tree) == 2:
self._query_tags.append(tag)
return tag
def untagged_term(self, tree):
"""Process an untagged term: sort (subterms)?."""
sort = tree[0]
if len(tree) > 1:
return {"tag": None, "sort": sort, "subterms": tree[1]}
return {"tag": None, "sort": sort, "subterms": None}
def unsorted_term(self, tree):
"""Process an unsorted term: TAG (subterms)?."""
tag = tree[0]
self.tags.add(tag)
if len(tree) > 1:
return {"tag": tag, "sort": None, "subterms": tree[1]}
return {"tag": tag, "sort": None, "subterms": None}
def query_term(self, tree):
"""Process a tagged term: TAG (":" untagged_term)."""
tag = tree[0]
sort = tree[1]
self.tags.add(tag)
if len(tree) > 2:
return {"tag": tag, "sort": sort, "subterms": tree[2]}
return {"tag": tag, "sort": sort, "subterms": None}
def q_subterm(self, tree):
"Process a subterm: FEATURE -> term."
feat = tree[0]
term = tree[1]
return (feat, term)
def q_subterms(self, tree):
"Process subterms: '(' subterm (',' subterm)* ')'."
d = defaultdict(list)
for k, v in tree:
d[k].append(v)
return d
def prefixed_query_term(self, tree):
return tree[-1]
def transform(self, parse_tree):
self._query_tags = list()
term = super().transform(parse_tree)
return term
[docs]
class QueryTermParser(OsfTermParser):
def __init__(self):
self.parser = Lark.open_from_package(
"fosf", FOSF_GRAMMAR, start="prefixed_query_term"
)
self.transformer = _QueryTermTransformer()
self.term_constructor = Term
self.tags = set()
self.tag_counter = 0
[docs]
def parse(
self, expression: str, default_tag="X", create_using=Term
) -> tuple[Term, dict]:
term = super().parse(expression, default_tag, create_using)
query_tags = self.transformer._query_tags
# Remove duplicate query tags if any
query_tags = list(dict.fromkeys(query_tags))
data = {
"vars": query_tags,
"prefixes": self.transformer.prefix,
"base": self.transformer.base,
}
return term, data
class _UnificationTransformer(_TaxonomyTransformer, _OsfTermTransformer):
def unif_program(self, tree):
return tree[0], tree[1], tree[2]
[docs]
class UnificationParser(OsfTermParser):
def __init__(self):
self.parser = Lark.open_from_package("fosf", FOSF_GRAMMAR, start="unif_program")
self.transformer = _UnificationTransformer()
self.tag_counter = 0
self.term_constructor = Term
self.tags = set()
@overload
def parse(
self, expression: str, default_tag="X", term_constructor=NormalTerm
) -> tuple[SortTaxonomy, NormalTerm, NormalTerm]: ...
@overload
def parse(
self, expression: str, default_tag="X", term_constructor=Term
) -> tuple[SortTaxonomy, Term, Term]: ...
[docs]
def parse(
self, expression: str, default_tag="X", term_constructor=Term
) -> tuple[SortTaxonomy, Term, Term]:
parse_tree = self.parser.parse(expression)
taxonomy, dict1, dict2 = self.transformer.transform(parse_tree)
self.tags = self.transformer.tags
self.tag_counter = 0
self.term_constructor = term_constructor
term1 = self._dict_to_term(dict1, default_tag)
term2 = self._dict_to_term(dict2, default_tag)
return taxonomy, term1, term2