phino-0.0.145: src/Rewriter.hs
{-# LANGUAGE DeriveAnyClass #-}
{-# LANGUAGE DerivingStrategies #-}
{-# LANGUAGE DuplicateRecordFields #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE OverloadedRecordDot #-}
{-# LANGUAGE RecordWildCards #-}
{-# OPTIONS_GHC -Wno-name-shadowing #-}
-- SPDX-FileCopyrightText: Copyright (c) 2025 Objectionary.com
-- SPDX-License-Identifier: MIT
module Rewriter (Seen, direct, every, fast, interpreted, rewrite, RewriteContext (..), Rewritten, Rewrittens, Rewrittens', seenInsert, seenMember, stepHeaders) where
import AST
import Builder
import Control.Exception (Exception, throwIO)
import Data.List.NonEmpty (NonEmpty (..))
import qualified Data.List.NonEmpty as NE
import qualified Data.Map.Strict as Map
import Data.Maybe (fromMaybe)
import Data.Set (Set)
import qualified Data.Set as Set
import Deps
import Locator (locatedExpression, withLocatedExpression)
import Logger (logDebug)
import Matcher (Subst, sites)
import Must (Must (..), exceedsUpperBound, inRange)
import Printer (printExpression)
import Replacer (ReplaceExpressionFunc, replaceExpression, replaceExpressionFast)
import Rule (RuleContext (RuleContext), Step (..))
import qualified Rule as R
import Text.Printf (printf)
import qualified Yaml as Y
type RewriteState = (NonEmpty Rewritten, Expression, Seen, Bool, Maybe (Set Int))
type Seen = Map.Map Int [Expression]
seenMember :: Int -> Expression -> Seen -> Bool
seenMember digest expr seen = maybe False (elem expr) (Map.lookup digest seen)
seenInsert :: Int -> Expression -> Seen -> Seen
seenInsert digest expr = Map.insertWith (++) digest [expr]
type Rewritten = (Expression, Maybe (Judgment, String))
type Rewrittens = (NonEmpty Rewritten, Bool)
type Rewrittens' = ([Rewritten], Bool)
stepHeaders :: [Rewritten] -> [String]
stepHeaders chain = zipWith3 header [1 ..] chain (Nothing : map Just chain)
where
header :: Int -> Rewritten -> Maybe Rewritten -> String
header step _ Nothing = printf "=== Step #%d" step
header step (current, _) (Just (before, rule)) =
printf
"=== Step #%d, Rule '%s', %dt -> %dt"
step
(maybe "?" snd rule)
(countNodes before)
(countNodes current)
type ToReplace = (Expression, Expression, Expression, [Subst])
data RewriteContext = RewriteContext
{ _locator :: Expression
, _maxDepth :: Int
, _maxCycles :: Int
, _depthSensitive :: Bool
, _universe :: Maybe Expression
, _buildTerm :: BuildTermFunc
, _normal :: Expression -> Bool
, _matching :: Maybe Expression -> Expression -> Set Int
, _must :: Must
, _breakpoint :: Maybe String
, _saveStep :: SaveStepFunc
}
data RewriteException
= MustBeGoing Must Int
| MustStopBefore Must Int
| StoppedOnLimit String Int
| LoopingRewriting String String Int
deriving (Exception)
instance Show RewriteException where
show (MustBeGoing mst cnt) =
printf
"With option --must=%s it's expected rewriting cycles to be in range [%s], but rewriting stopped after %d cycles"
(show mst)
(show mst)
cnt
show (MustStopBefore mst cnt) =
printf
"With option --must=%s it's expected rewriting cycles to be in range [%s], but rewriting has already reached %d cycles and is still going"
(show mst)
(show mst)
cnt
show (StoppedOnLimit flg lim) =
printf
"With option --depth-sensitive it's expected rewriting iterations amount does not reach the limit: --%s=%d"
flg
lim
show (LoopingRewriting expr rul stp) =
printf
"On rewriting step '%d' of rule '%s' we got the same expression as we got at one of the previous step, it seems rewriting is looping\nExpression: %s"
stp
rul
expr
buildAndReplace' :: ToReplace -> ReplaceExpressionFunc -> IO Expression
buildAndReplace' (expr, ptn, res, substs) func = do
ptns <- buildExpressionsThrows ptn substs
repls <- buildExpressionsThrows res substs
pure (func (expr, ptns, map const repls))
tryBuildAndReplaceFast :: ToReplace -> IO Expression
tryBuildAndReplaceFast state@(expr, ptn@(ExFormation (_ : pbds)), res@(ExFormation (_ : rbds)), substs)
| fast ptn res = do
logDebug "Applying fast replacing since 'pattern' and 'result' are suitable for this..."
buildAndReplace' (expr, ExFormation (init pbds), ExFormation (init rbds), substs) replaceExpressionFast
| otherwise = do
logDebug "Applying regular replacing..."
buildAndReplace' state replaceExpression
tryBuildAndReplaceFast state = buildAndReplace' state replaceExpression
fast :: Expression -> Expression -> Bool
fast (ExFormation _pbds@(pbd : pbds)) (ExFormation _rbds@(rbd : rbds)) =
startsAndEndsWithMeta _pbds
&& startsAndEndsWithMeta _rbds
&& pbd == rbd
&& last pbds == last rbds
&& not (hasMetaBindings (init pbds))
&& not (hasMetaBindings (init rbds))
where
startsAndEndsWithMeta :: [Binding] -> Bool
startsAndEndsWithMeta [] = False
startsAndEndsWithMeta bds@(bd : _) =
length bds > 1
&& isMetaBinding bd
&& isMetaBinding (last bds)
hasMetaBindings :: [Binding] -> Bool
hasMetaBindings = foldl (\acc bd -> acc || isMetaBinding bd) False
isMetaBinding :: Binding -> Bool
isMetaBinding = \case
BiMeta _ -> True
BiAny _ -> True
_ -> False
fast _ _ = False
interpreted :: Y.Rule -> Step
interpreted rule = Step rule.name applied
where
applied :: RuleContext -> Expression -> IO (Maybe Expression)
applied ctx expr =
R.matchExpressionWithRule expr rule ctx >>= \case
[] -> pure Nothing
matched -> Just <$> tryBuildAndReplaceFast (expr, rule.pattern, rule.result, matched)
direct :: String -> Bool -> (Maybe Expression -> Expression -> [Expression]) -> Step
direct name redex rewritten = Step name applied
where
applied :: RuleContext -> Expression -> IO (Maybe Expression)
applied (RuleContext _ universe _) expr = pure $ case sites redex (rewritten universe) expr of
[] -> Nothing
found -> Just (replaceExpression (expr, map fst found, map (const . snd) found))
every :: [Step] -> Maybe Expression -> Expression -> Set Int
every steps _ _ = Set.fromList (zipWith const [0 ..] steps)
rewrite' :: RewriteState -> [(Int, Step)] -> Int -> RewriteContext -> IO RewriteState
rewrite' state [] _ _ = pure state
rewrite' (rewrittens, located, unique, stop, found) ((idx, rule) : rest) iteration ctx@RewriteContext{..}
| Set.member idx matched || _breakpoint == Just (_name rule) =
_rewrite (rewrittens, located, unique, stop, Just matched) 1 >>= \case
state'@(_, _, _, True, _) -> pure state'
state' -> rewrite' state' rest iteration ctx
| otherwise = rewrite' (rewrittens, located, unique, stop, Just matched) rest iteration ctx
where
matched :: Set Int
matched = fromMaybe (_matching _universe located) found
_rewrite :: RewriteState -> Int -> IO RewriteState
_rewrite (_rewrittens@((current, _) :| _), expression, _unique, _, _found) _count =
let ruleName = _name rule
in if _count - 1 == _maxDepth
then do
logDebug (printf "Max amount of rewriting cycles (%d) for rule '%s' has been reached, rewriting is stopped" _maxDepth ruleName)
if _depthSensitive
then do
exhausted <- applicable expression [rule] ctx
if exhausted
then throwIO (StoppedOnLimit "max-depth" _maxDepth)
else pure (_rewrittens, expression, _unique, False, _found)
else pure (_rewrittens, expression, _unique, False, _found)
else do
logDebug (printf "Starting rewriting cycle for rule '%s': %d out of %d" ruleName _count _maxDepth)
_applied rule (RuleContext _buildTerm _universe _normal) expression >>= \case
Nothing -> do
logDebug (printf "Rule '%s' does not match, rewriting is stopped" ruleName)
if _breakpoint == Just ruleName
then do
logDebug (printf "Rule '%s' is a breakpoint, dropping down all the previous rewritings..." ruleName)
pure (_rewrittens, expression, _unique, True, _found)
else pure (_rewrittens, expression, _unique, False, _found)
Just expr -> do
logDebug (printf "Rule '%s' has been matched and applied" ruleName)
if expression == expr
then do
logDebug (printf "Applied '%s', no changes made" ruleName)
pure (_rewrittens, expression, _unique, False, _found)
else
let digest = hashExpression expr
in if seenMember digest expr _unique
then throwIO (LoopingRewriting (printExpression expr) ruleName _count)
else do
logDebug
( printf
"Applied '%s' (%d nodes -> %d nodes)\n%s"
ruleName
(countNodes expression)
(countNodes expr)
(printExpression expr)
)
updated <- withLocatedExpression _locator expr current
_saveStep updated
_rewrite (leadsTo updated, expr, seenInsert digest expr _unique, False, Nothing) (_count + 1)
where
leadsTo :: Expression -> NonEmpty Rewritten
leadsTo next =
let (head', _) :| rest = _rewrittens
in (next, Nothing) :| (head', Just (Normalization, _name rule)) : rest
applicable :: Expression -> [Step] -> RewriteContext -> IO Bool
applicable _ [] _ = pure False
applicable expression (rule : rest) ctx@RewriteContext{..} =
_applied rule (RuleContext _buildTerm _universe _normal) expression >>= \case
Nothing -> applicable expression rest ctx
Just _ -> pure True
rewrite :: Expression -> [Step] -> RewriteContext -> IO Rewrittens
rewrite expr rules ctx@RewriteContext{..} = do
located <- locatedExpression _locator expr
(rewrittens, exceeded) <- _rewrite ((expr, Nothing) :| [], located, Map.empty, False, Nothing) 0
pure (NE.reverse rewrittens, exceeded)
where
_rewrite :: RewriteState -> Int -> IO Rewrittens
_rewrite state@(rewrittens@((current, _) :| _), expression, _, _, _) count
| not (inRange _must count) && count > 0 && exceedsUpperBound _must count = throwIO (MustStopBefore _must count)
| count == _maxCycles && not (inRange _must count) = throwIO (MustBeGoing _must count)
| count == _maxCycles = do
logDebug (printf "Max amount of rewriting cycles for all rules (%d) has been reached, rewriting is stopped" _maxCycles)
if _depthSensitive
then do
exhausted <- applicable expression rules ctx
if exhausted
then throwIO (StoppedOnLimit "max-cycles" _maxCycles)
else pure (rewrittens, False)
else pure (rewrittens, True)
| otherwise = do
logDebug (printf "Starting rewriting cycle for all rules: %d out of %d" count _maxCycles)
rewrite' state (zip [0 ..] rules) count ctx >>= \case
(_, _, _, True, _) -> pure ((expr, Nothing) :| [], False)
state'@(rewrittens'@((current', _) :| _), _, _, False, _) ->
if length rewrittens' == length rewrittens || current' == current
then do
logDebug "Rewriting is stopped since it has no effect"
if not (inRange _must count)
then throwIO (MustBeGoing _must count)
else pure (rewrittens', False)
else _rewrite state' (count + 1)