module Tessera.Tree where import Data.List (foldl') import Data.Maybe (mapMaybe) import Tessera.Env import Tessera.Matrix import Tessera.Syntax data Occurrence = Occurrence [Int] deriving (Eq, Show) data Tree = Success Int | Failure | Switch Occurrence [(Key, Tree)] (Maybe Tree) | Bind Occurrence String Tree deriving (Eq, Show) data Strategy = Leftmost | Heuristic deriving (Eq, Show) strategyName :: Strategy -> String strategyName Leftmost = "leftmost" strategyName Heuristic = "heuristic" removeAt :: Int -> [a] -> [a] removeAt index items = take index items ++ drop (index + 1) items type Row = ([Pattern], Int) specializeRows :: Env -> Key -> [Row] -> [Row] specializeRows env key = concatMap step where arity = keyArity env key step (row, index) = case row of [] -> [] (p : rest) | isVariable p -> [(replicate arity PWild ++ rest, index)] | headKey p == Just key -> [(patternArguments p ++ rest, index)] | otherwise -> [] defaultRows :: [Row] -> [Row] defaultRows = mapMaybe step where step (row, index) = case row of (p : rest) | isVariable p -> Just (rest, index) _ -> Nothing compileMatch :: Env -> Strategy -> [Type] -> [Row] -> [Occurrence] -> Tree compileMatch env strategy types rows occurrences | null rows = Failure | all isVariable (fst (head rows)) = let (row, index) = head rows in foldr bindVariable (Success index) (zip row occurrences) | otherwise = let patterns = map fst rows column = chooseColumn strategy patterns columnType = types !! column occurrence = occurrences !! column restTypes = removeAt column types restOccurrences = removeAt column occurrences columnPatterns = map (!! column) patterns appearing = mapMaybe headKey columnPatterns path = occurrencePath occurrence in case signature env columnType of IntegerSignature -> let literals = [n | KeyInt n <- appearing] branches = [ (KeyInt n, compileMatch env strategy restTypes (specializeRows env (KeyInt n) rows) restOccurrences) | n <- literals ] defaultBranch = if any isVariable columnPatterns then Just (compileMatch env strategy restTypes (defaultRows rows) restOccurrences) else Nothing in Switch occurrence branches defaultBranch FiniteSignature constructors -> let keys = map fst constructors complete = allConstructorsPresent constructors columnPatterns branchKeys = if complete then keys else [k | k <- keys, k `elem` appearing] branches = [ ( key , compileMatch env strategy (argumentTypes ++ restTypes) (specializeRows env key rows) (argumentOccurrences ++ restOccurrences) ) | key <- branchKeys , let argumentTypes = maybe [] id (lookup key constructors) , let argumentOccurrences = [Occurrence (path ++ [j]) | j <- [0 .. length argumentTypes - 1]] ] defaultBranch = if not complete && not (null (defaultRows rows)) then Just (compileMatch env strategy restTypes (defaultRows rows) restOccurrences) else Nothing in Switch occurrence branches defaultBranch where bindVariable (pattern', occurrence) inner = case pattern' of PVar name -> Bind occurrence name inner _ -> inner occurrencePath :: Occurrence -> [Int] occurrencePath (Occurrence path) = path chooseColumn :: Strategy -> [[Pattern]] -> Int chooseColumn Leftmost matrix = head candidates where candidates = [i | i <- [0 .. width - 1], not (isVariable (head matrix !! i))] width = length (head matrix) chooseColumn Heuristic matrix = foldl' better (head candidates) (tail candidates) where width = length (head matrix) candidates = [i | i <- [0 .. width - 1], not (isVariable (head matrix !! i))] better best candidate = let score i = (distinctKeys i, negate (constructorCount i), i) in if score candidate < score best then candidate else best distinctKeys i = length (unique [k | Just k <- map (headKey . (!! i)) matrix]) constructorCount i = length [() | row <- matrix, not (isVariable (row !! i))] unique = foldl' (\acc x -> if x `elem` acc then acc else acc ++ [x]) [] initialOccurrences :: Int -> [Occurrence] initialOccurrences n = [Occurrence [i] | i <- [0 .. n - 1]] treeSize :: Tree -> Int treeSize tree = case tree of Success _ -> 1 Failure -> 1 Bind _ _ inner -> 1 + treeSize inner Switch _ branches fallback -> 1 + sum (map (treeSize . snd) branches) + maybe 0 treeSize fallback treeDepth :: Tree -> Int treeDepth tree = case tree of Success _ -> 1 Failure -> 1 Bind _ _ inner -> 1 + treeDepth inner Switch _ branches fallback -> let branchDepth = maximum (0 : map (treeDepth . snd) branches) fallbackDepth = maybe 0 treeDepth fallback in 1 + max branchDepth fallbackDepth