-- Algorithms in Haskell for
-- @article{APD05,
-- 	TITLE = "Automated Pattern Detection: An Algorithm for Constructing
--	Optimally Synchronizing Multi-Regular Language Filters",
--	AUTHOR = "C. S. McTague and J. P. Crutchfield",
--	JOURNAL = "Theoretical Computer Science",
--	NOTE = "In press"
--	}
-- Copyright (2005) McTague and Crutchfield

module APD where
import List
import Maybe
import Array
data FA s i = FA { faStarts :: [s], faTrans :: [(s,i,s)], faFinals :: [s] }
            deriving (Show, Read, Eq)
type Transducer s i o = FA s (i,o)
faAlphabet :: Eq i => FA s i -> [i]
faAlphabet fa = nub [ a | (_,a,_) <- faTrans fa ]

faStates :: Eq s => FA s i -> [s]
faStates fa = faStarts fa `union` (transStates $ faTrans fa) `union` faFinals fa

transStates :: Eq s => [(s,i,s)] -> [s]
transStates trans = nub $ foldl ( \ss (s,_,s')->s:s':ss ) [] trans
faDet :: (Ord s, Eq i) => FA s i -> FA [s] i
faDet fa = FA starts' trans' finals'
    where starts' = [sort $ faStarts fa]
	  trans' = faTransFromDelta starts' delta $ faAlphabet fa
          delta detState a | detState'==[] = Nothing
                           | otherwise = Just detState'
              where detState' = sort $ nub [ s'' | s<-detState,
                                                   (s',a',s'')<-faTrans fa, s'==s, a'==a ]
          finals' = filter containsFinal (starts' `union` transStates trans')
          containsFinal detState = or [ s `elem` faFinals fa | s <- detState ]

faIntersect :: (Eq s, Eq s', Eq i) => --assumes $\Varid{fa}$ and $\Varid{fa'}$ are semi-deterministic
                  FA s i -> FA s' i -> FA (s,s') i
faIntersect fa fa' =
    FA starts'' trans'' finals''
    where starts'' = [ (s,s') | s<-faStarts fa, s'<-faStarts fa' ]
          alphabet = faAlphabet fa `union` faAlphabet fa'
          trans'' = faTransFromDelta starts'' delta alphabet
          delta (s,s') a | null ss || null ss' = Nothing
                         | otherwise = Just (head ss, head ss') 
              where ss  = [ s''' | (s'',a',s''')<-faTrans fa, s''==s, a==a' ]
                    ss' = [ s''' | (s'',a',s''')<-faTrans fa', s''==s', a==a' ]
          finals'' = filter bothFinal $ union starts'' $ transStates trans''
          bothFinal (s,s') = s `elem` faFinals fa && s' `elem` faFinals fa'

faTransFromDelta :: (Eq s, Eq i) => [s]->(s->i->Maybe s)->[i]->[(s,i,s)]
faTransFromDelta starts delta alphabet = trans
        where f states | null (states' \\ states) = states
                       | otherwise = f $ nub $ states++states'
                       where states' = nub $ catMaybes [ delta s i | s <- states, i<-alphabet ]
              states = f starts
              trans = [ (s,i,fromJust s') | s<-states, i<-alphabet, let s'=delta s i, isJust s' ]
faDisjointUnion :: [FA s i] -> FA (Int,s) i
faDisjointUnion fas = concatFas $ zipWith renameStates [1..] fas
    where renameStates n fa = FA starts' trans' finals'
              where starts' = [ (n,s) | s <- faStarts fa ]
                    trans' = [ ( (n,s), a, (n,s') ) | (s,a,s') <- faTrans fa ]
                    finals' = [ (n,s) | s <- faFinals fa ]
          concatFas [] = FA [] [] []
          concatFas (FA starts trans finals:fas) = 
              FA (starts++starts') (trans++trans') (finals++finals')
              where FA starts' trans' finals' = concatFas fas
transducerFilterFromDomains :: (Eq s, Ord s, Eq i) => [FA s i] -> Transducer [(Int,s)] i Int
transducerFilterFromDomains faDs =
    FA (faStarts faA) (baseTTrans++newTTrans) (faFinals faA)
    where faA = faDet $ faDisjointUnion faDs -- $\mc{A}:=\Det(\mc{D}_1 \sqcup \cdots \sqcup \mc{D}_n)$
          baseTTrans = [ (s,(a,f s'),s') | (s,a,s')<-faTrans faA ]
              where f ss | length is == 1 = head is -- transition ending in $\mc{D}_i$
			 | otherwise = 0 -- synchronization, $\lambda$
			 where is = nub $ map fst ss
          forbiddenPairs = [ (s,a) | s<-faStates faA, a<-faAlphabet faA ] \\ [ (s,a) | (s,a,_)<-faTrans faA ]
          newTTrans = map newTransition forbiddenPairs
          newTransition (s,a) = (s,(a,o),s')
              where faAsa = (FA (f:faStates faA) ((s,a,f):faTrans faA) [f]) -- the automaton $\mc{A}^{s,a}$
                    f = [(1+length faDs, head $ faStarts $ head faDs)] -- the fresh state $f$ used to build $\mc{A}^{s,a}$
                    faDetAsaCapA = (faDet faAsa) `faIntersect` faA -- the automaton $\Det(\mc{A}^{s,a} \cap \mc{A})$
                    faZ = faDet $ zero faDetAsaCapA -- the automaton $\Det(\mc{Z}[\Det(\mc{A}^{s,a} \cap \mc{A})])$
                    reachableStateSeq = take (length $ faStates faZ) -- the states $s_1, \dots, s_{m+m'}$
                                        (iterate nextState $ head $ faStarts faZ)
                        where nextState s = head [ s'' | (s',_,s'')<-faTrans faZ, s'==s ] 
                    sStarLs = map ((map snd) . -- the sets $\{S_{*,\ell}\}_{\ell=1}^{m+m'}$
                                   (intersect $ faFinals faDetAsaCapA)) reachableStateSeq
                    sdStars = [ nub -- the sets $\{S_{d,*}\}_{i=1}^{n}$
                                [ s | s<-map snd $ faFinals faDetAsaCapA, d==length (nub (map fst s)) ]
                              | d <- [1..length faDs] ]
                    sdls = concatMap (\sdStar->map (intersect sdStar) sStarLs) sdStars -- the sets $\{S_{d,\ell}\}$
                    [s'] = head $ filter (\sdl->length sdl==1) sdls -- the state $s'$ to which to synchronize
                    o = -1 -- domain break
          zero fa = FA (faStarts fa) [(s,0,s') | (s,_,s')<-faTrans fa] (faFinals fa)
faDisjoint :: (Ord s, Ord s', Eq i) => FA s i -> FA s' i -> Bool
faDisjoint fa fa' = faNull $ fa `faIntersect` fa'
faNull :: (Ord s, Eq i) => FA s i -> Bool
faNull fa | null $ faFinals faD = True
          | not $ null (faFinals faD \\ faStarts faD) = False
          | not $ null [ s | (s,_,s') <- faTrans faD, s' `elem` faStarts faD ] = False
          | otherwise = True
    where faD = faDet fa
faDifference :: (Eq s, Eq i, Ord s') => FA s i -> FA s' i -> FA (s,[s']) i
faDifference fa fa' = fa `faIntersect` faComplement alphabet fa'
    where alphabet = faAlphabet fa `union` faAlphabet fa'
faComplement :: (Ord s, Eq i) => [i] -> FA s i -> FA [s] i
faComplement alphabet fa = FA starts trans finals
    where faD = faDet fa
          starts = faStarts faD
	  trans = (faTrans faD)
		  ++ [ (s,a,[]) | (s,a) <- forbiddenPairs alphabet faD, not $ null s ]
		  ++ [ ([],a,[]) | a <- alphabet ]
          finals = ([] : faStates faD) \\ faFinals faD

forbiddenPairs :: (Eq s, Eq i) => [i] -> FA s i -> [(s,i)]
forbiddenPairs alphabet fa = 
    [ (s,a) | s<-faStates fa, a<-alphabet ] \\ [ (s,a) | (s,a,_)<-faTrans fa ]
faDisjoin :: (Ord s, Eq i) => [FA s i] -> [FA Int i]
faDisjoin [] = []
faDisjoin (fa:fas) = filter (not.faNull) (faMinusRest : disjoinedRestMinusFa)
    where faMinusRest = faIS $ fa `faDifference` (faIS $ faDisjointUnion fas)
          disjoinedRestMinusFa = concatMap f (faDisjoin fas)
          f fa' | fa `faDisjoint` fa' = [fa']
		| otherwise = [faIS $ fa `faIntersect` fa',
			       faIS $ fa' `faDifference` fa]
faIntegerizeStates, faIS :: Eq s => FA s i -> FA Int i
faIS = faIntegerizeStates
faIntegerizeStates fa =
    FA starts trans finals
    where table = zip (faStates fa) [1..]
          starts = catMaybes [ lookup s table | s <- faStarts fa ]
          trans = [ (i,a,j) | (s,a,s') <- faTrans fa, 
                    let Just i = lookup s  table,
                    let Just j = lookup s' table ]
          finals = catMaybes [ lookup s table | s <- faFinals fa ]
data Epsilon a = Epsilon | Letter a
		 deriving Show

faRemoveEpsilons :: Eq s => FA s (Epsilon i) -> FA s i
faRemoveEpsilons fa = FA starts trans finals
    where starts = epsilonClosure fa $ faStarts fa
          finals = reverseEpsilonClosure fa $ faFinals fa
	  trans = concatMap inducedTrans $ faTrans fa
	  inducedTrans (_,Epsilon,_) = []
	  inducedTrans (s,Letter a,s') =
              [ (s'',a,s''') | s'' <- reverseEpsilonClosure fa [s],
                               s''' <- reverseEpsilonClosure fa [s'] ]

epsilonClosure, reverseEpsilonClosure :: Eq s => FA s (Epsilon i) -> [s] -> [s]
epsilonClosure fa states = fixedPointBy (==) $ iterate epsilonAway states
    where epsilonAway ss = nub $ ss ++ [ s' | (s,Epsilon,s') <- faTrans fa, s `elem` ss ]
reverseEpsilonClosure = epsilonClosure . faReverse

faReverse :: FA s i -> FA s i
faReverse fa = FA (faFinals fa) trans (faStarts fa)
    where trans = [ (s',a,s) | (s,a,s') <- faTrans fa ]

fixedPointBy :: (a->a->Bool) -> [a] -> a --Assumes that [s] has a fixed point
fixedPointBy eqFunc (a1:a2:as) | a1 `eqFunc` a2 = a1
			       | otherwise = fixedPointBy eqFunc (a2:as)

faConcat :: (Ord s, Eq i) => FA s i -> FA s i -> FA [(Int,s)] i
faConcat fa fa' = faDet $ faRemoveEpsilons (FA starts trans finals)
    where starts = [ (1,s) | s <- faStarts fa ]
          finals = [ (2,s) | s <- faFinals fa' ]
          trans = [ ((1,s),Letter i,(1,s')) | (s,i,s') <- faTrans fa ]
		  ++ [ ((1,s),Epsilon,(2,s')) | s <- faFinals fa, s' <- faStarts fa' ]
                  ++ [ ((2,s),Letter i,(2,s')) | (s,i,s') <- faTrans fa' ]
initialGamma :: (Ord s, Eq i) => [FA s i] -> [[(s,[FA Int i])]]
initialGamma faDs =
    [ [ (s, gamma [(i,s)]) | s <- faStates (faDs!!(i-1)) ] | i <- [1..length faDs] ]
    where faA = faDet $ faDisjointUnion faDs
          alphabet = faAlphabet faA
          gamma s | null forbiddenAs = [faSigmaStar alphabet]
                  | otherwise = map (faIS.faMin) $
                      faDisjoin [ faSigmaStar alphabet `faConcat` faBsas'
                                            | a <- forbiddenAs, faBsas' <- gamma' s a ]
                  --let $\Gamma(s) := \mrm{Disjoin}( \{ \Sigma^* \cdot \mc{B}(s,a,s') : (s,a,s'') \not \in T(\mc{A}), s' \in S(\mc{A}) \})$
              where forbiddenAs = alphabet \\ [ a | (s',a,_) <- faTrans faA, s'==s ]
          gamma' s a = [ faIS $ FA (faStarts faDetAsaCapA) (faTrans faDetAsaCapA)
                                [ s'' | (s'',_,(_,s''')) <- faTrans faDetAsaCapA, s'''==s' ]
                                    | s' <- resyncStates ]
              where faAsa = (FA (f:faStates faA) ((s,a,f):faTrans faA) [f]) --the automaton $\mc{A}^{s,a}$
 		    f = [(1+length faDs, head $ faStarts $ head faDs)] --the fresh state used to build $\mc{A}^{s,a}$
                    faDetAsaCapA = (faDet faAsa) `faIntersect` faA --the automaton $\Det(\mc{A}^{s,a}) \cap \mc{A}$
                    resyncStates = nub [ s' | (_,s') <- faFinals faDetAsaCapA ] --the states $s'$ that actually matter

faPullBackFinalsBy :: (Eq s, Eq i) => i -> FA s i -> FA s i
faPullBackFinalsBy a fa = FA (faStarts fa) (faTrans fa) finals
    where finals = [ s | (s,a',s') <- faTrans fa, a'==a, s' `elem` faFinals fa ]

faSigmaStar :: [i] -> FA Int i
faSigmaStar alphabet = FA [1] [ (1,a,1) | a<-alphabet ] [1]
faMin :: (Ord s, Eq i) => FA s i -> FA [[s]] i
faMin = faDet . faReverse . faDet . faReverse
refineGammaForSingleD :: (Ix s, Eq i) => FA s i -> [(s,[FA Int i])] -> [(s,[FA Int i])]
refineGammaForSingleD faD associations =
    assocs $ foldl refineArrayWithTransition initialGammaArray $ faTrans faD
    where initialGammaArray = array arrayBounds associations
          arrayBounds = (minimum $ faStates faD, maximum $ faStates faD)
          refineArrayWithTransition array (s,a,s') = accum (curry snd) array [(s,gamma')]
              where gamma' = map (faIS.faMin) $ filter (not.faNull)
                             [ faPullBackFinalsBy a $ faIS $
                               (faE `faConcat` FA [1] [(1,a,2)] [2]) `faIntersect` faE'
                                   | faE <- array!s, faE' <- array!s' ]

optimalGamma :: (Ix s, Eq i) => [FA s i] -> [[(s,[FA Int i])]]
optimalGamma faDs =
    fixedPointBy cardinalityEq $ iterate (zipWith refineGammaForSingleD faDs) $ initialGamma faDs
    where cardinality aa = sum [ length gammaOfs | associations <- aa, (_,gammaOfs) <- associations ]
          cardinalityEq aa aa' = cardinality aa == cardinality aa'
optimizeDomains :: (Ix s, Ord s, Eq i) =>  [FA s i] -> [FA Int i]
optimizeDomains faDs =
    zipWith optimalDomainFromGamma faDs $ optimalGamma faDs
    where optimalDomainFromGamma faD associations = faIS $ FA starts trans finals
              where starts = [ (s,faE) | s <- faStarts faD, faE <- gamma s ]
                    trans = [ ( (s,faE), a, (s',faE') ) | (s,a,s')<-faTrans faD,
			       faE <- gamma s, faE' <- gamma s',
			       (faIS (faE `faConcat` (FA [1] [(1,a,2)] [2]))) `faSubset` faE' ]
                    finals = [ (s,faE) | s <- faFinals faD, faE <- gamma s ]
                    gamma s = fromJust $ lookup s associations

faSubset :: (Ord s, Eq i) => FA s i -> FA s i -> Bool
faSubset fa fa' = faNull $ fa `faDifference` fa'
fig7Domains = [FA [1,2] [(1,0,2),(1,1,2),(2,0,1)] [1,2],
               FA [3..6] [(3,1,4),(4,1,5),(5,0,6),(6,0,3),(6,1,3)] [3..6]]
