strict-impl-params 1.1.1 → 1.1.2
raw patch · 4 files changed
+241/−7 lines, 4 filesdep +strict-impl-paramsdep ~ghcPVP ok
version bump matches the API change (PVP)
Dependencies added: strict-impl-params
Dependency ranges changed: ghc
API changes (from Hackage documentation)
Files
- StrictImplParams.hs +32/−5
- strict-impl-params.cabal +17/−2
- test/Defs.hs +99/−0
- test/Main.hs +93/−0
StrictImplParams.hs view
@@ -6,6 +6,7 @@ import Data.Foldable import Data.Maybe+import GHC.Core.Class (classMethods) import GHC.Core.Predicate import GHC.Plugins @@ -54,10 +55,25 @@ isIPPred_maybe cls tys #endif --- | Force var, continue with CoreExpr body that has Type type.+-- | The value stored in an implicit parameter dictionary, together with its type,+-- i.e. @(ip \@x \@a d, a)@ for @d :: IP x a@.+ipValue :: Var -> Maybe (CoreExpr, Type)+ipValue x = do+ (tc, args) <- splitTyConApp_maybe (varType x)+ cls <- tyConClass_maybe tc+ valTy <- case args of [_sym, valTy] -> Just valTy; _ -> Nothing+ sel <- case classMethods cls of [sel] -> Just sel; _ -> Nothing+ pure (mkCoreApps (Var sel) (foldr' (\a as -> (Type a:) $! as) [Var x] args), valTy)++-- | Force the value of an implicit param, continue with CoreExpr body that has Type+-- type.+--+-- Note: we force the *value* stored in the dictionary, by applying the class method+-- to it. forceVar :: Var -> CoreExpr -> Type -> CoreExpr-forceVar x body bodyTy =- mkWildCase (Var x) (GHC.Scaled manyType (varType x)) bodyTy [Alt DEFAULT [] body]+forceVar x body bodyTy = case ipValue x of+ Just (val, valTy) -> mkWildCase val (GHC.Scaled manyType valTy) bodyTy [Alt DEFAULT [] body]+ Nothing -> error "forceVar: not an implicit parameter dictionary" pass :: ModGuts -> CoreM ModGuts pass guts = do@@ -72,8 +88,19 @@ go :: [Var] -> CoreExpr -> Type -> CoreExpr go vars t a = case t of Lam x t -> case coreFullView a of- GHC.ForAllTy _ a -> Lam x $! go vars t a- GHC.FunTy _ _ a b | isImplicitParamTy a -> Lam x $! go (x:vars) t b+ -- Note: a Core type lambda binds its own type variable, which is not necessarily the+ -- one bound by the ForAllTy in the definition's type. Since we use the def type to+ -- generate the Core of the forcing, we need to rename the type binder to match the+ -- lambda binder. We sincerely hope that Core names are unique and this renaming doesn't+ -- introduce shadowing.+ GHC.ForAllTy bndr a -> let tv = binderVar bndr+ a' | isTyVar x, tv /= x = substTyWith [tv] [mkTyVarTy x] a+ | otherwise = a+ in Lam x $! go vars t a'+ -- Precaution: we add an occurrence of the binder, so whatever occurrence info it+ -- carries from the desugarer should be outdated.+ GHC.FunTy _ _ a b | isImplicitParamTy a -> let x' = zapIdOccInfo x in+ Lam x' $! go (x':vars) t b | otherwise -> Lam x $! go vars t b _ -> error $ "unexpected type for lambda expression: " ++ dbg a t ->
strict-impl-params.cabal view
@@ -1,7 +1,7 @@ name: strict-impl-params cabal-version: 1.12 build-type: Simple-version: 1.1.1+version: 1.1.2 synopsis: Plugin for making implicit parameters strict homepage: https://github.com/AndrasKovacs/ghc-strict-implicit-params bug-reports: https://github.com/AndrasKovacs/ghc-strict-implicit-params/issues@@ -15,5 +15,20 @@ default-language: Haskell2010 build-depends: base >=4.7 && <5- , ghc >= 9.2.3 && < 9.16+ , ghc >= 9.4.8 && < 9.16 exposed-modules: StrictImplParams++test-suite strictness+ default-language: Haskell2010+ type: exitcode-stdio-1.0+ hs-source-dirs: test+ main-is: Main.hs+ other-modules: Defs+ build-depends:+ base >=4.7 && <5+ , strict-impl-params+ -- -O2 and -dcore-lint needed to catch GHC+ -- miscompilation, but Hackage does not accept+ -- -dcore-lint+ ghc-options: -O2 -fplugin StrictImplParams+ -- ghc-options: -dcore-lint
+ test/Defs.hs view
@@ -0,0 +1,99 @@+{-# language ImplicitParams, RankNTypes #-}++module Defs where++f :: (?x :: Int) => Int -> Int+f n = g n + g (n + 1)+{-# noinline f #-}++g :: (?x :: Int) => Int -> Int+g n = n + ?x+{-# noinline g #-}++h :: (?x :: Int) => (?y :: Int) => Int -> Int+h n = f n + ?y+{-# noinline h #-}++dead :: (?x :: Int) => Int -> Int+dead n = n + 1+{-# noinline dead #-}++deadPass :: (?x :: Int) => Int -> Int+deadPass n = dead n+{-# noinline deadPass #-}++rec1 :: (?x :: Int) => Int -> Int+rec1 0 = ?x+rec1 n = rec1 (n - 1)+{-# noinline rec1 #-}++branch :: (?x :: Int) => Bool -> Int -> Int+branch True n = n + ?x+branch False n = n+{-# noinline branch #-}++local :: (?x :: Int) => Int -> Int+local n = go n+ where go k = k * ?x+{-# noinline local #-}++localIP :: Int -> Int+localIP n = let go :: (?z :: Int) => Int -> Int+ go k = k + ?z+ {-# noinline go #-}+ in let ?z = 7 in go n + go (n + 1)+{-# noinline localIP #-}++poly :: (?x :: Int) => a -> (a, Int)+poly a = (a, g 0)+{-# noinline poly #-}++mixed :: (?x :: Int) => forall a. (a -> Int) -> a -> Int+mixed k a = k a + ?x+{-# noinline mixed #-}++str :: (?s :: String) => Int -> String+str n = show n ++ ?s+{-# noinline str #-}++strTwice :: (?s :: String) => String+strTwice = str 1 ++ str 2+{-# noinline strTwice #-}++data Cfg = Cfg { cfgA :: !Int, cfgB :: String } deriving (Eq, Show)++cfgUse :: (?cfg :: Cfg) => Int -> String+cfgUse n = cfgB ?cfg ++ show (cfgA ?cfg + n)+{-# noinline cfgUse #-}++cfgPass :: (?cfg :: Cfg) => Int -> String+cfgPass n = cfgUse n ++ cfgUse (n + 1) ++ show ?cfg+{-# noinline cfgPass #-}++fnUse :: (?k :: Int -> Int) => Int -> Int+fnUse n = ?k n+{-# noinline fnUse #-}++fnPass :: (?k :: Int -> Int) => Int -> Int+fnPass n = fnUse n + fnUse (n + 1)+{-# noinline fnPass #-}++c0 :: (?x :: Int) => Int -> Int+c0 n = n + ?x+{-# noinline c0 #-}++c1 :: (?x :: Int) => Int -> Int+c1 n = c0 n+{-# noinline c1 #-}++c2 :: (?x :: Int) => Int -> Int+c2 n = c1 n+{-# noinline c2 #-}++c3 :: (?x :: Int) => Int -> Int+c3 n = c2 n + c1 n + c0 n+{-# noinline c3 #-}++closure :: (?x :: Int) => Int -> (Int -> Int)+closure n = \k -> k + n + ?x+{-# noinline closure #-}
+ test/Main.hs view
@@ -0,0 +1,93 @@+{-# language ImplicitParams #-}++module Main (main) where++import Control.Exception+import Data.IORef+import System.Exit++import Defs++main :: IO ()+main = do+ ref <- newIORef (0 :: Int)++ -- Correct results.+ let ?x = 10 in do+ eq ref "g" (g 5) 15+ eq ref "f" (f 5) 31+ eq ref "dead" (dead 5) 6+ eq ref "deadPass" (deadPass 5) 6+ eq ref "rec1" (rec1 4) 10+ eq ref "branchT" (branch True 5) 15+ eq ref "branchF" (branch False 5) 5+ eq ref "local" (local 5) 50+ eq ref "poly" (poly 'a') ('a', 10)+ eq ref "mixed" (mixed length "abc") 13+ eq ref "c3" (c3 1) 33+ eq ref "closure" (closure 1 2) 13+ let ?y = 100 in eq ref "h" (h 5) 131++ eq ref "localIP" (localIP 1) 17++ let ?s = "!" in do+ eq ref "str" (str 7) "7!"+ eq ref "strTwice" strTwice "1!2!"++ let ?cfg = Cfg 3 "c" in do+ eq ref "cfgUse" (cfgUse 1) "c4"+ eq ref "cfgPass" (cfgPass 1) ("c4" ++ "c5" ++ show (Cfg 3 "c"))++ let ?k = (* 2) in do+ eq ref "fnUse" (fnUse 3) 6+ eq ref "fnPass" (fnPass 3) 14++ -- Strictness: each of these must force the implicit parameter.+ let boom = error "boom" :: Int+ strict ref "g" (\_ -> let ?x = boom in g 1)+ strict ref "f" (\_ -> let ?x = boom in f 1)+ strict ref "dead" (\_ -> let ?x = boom in dead 1)+ strict ref "deadPass" (\_ -> let ?x = boom in deadPass 1)+ strict ref "rec1" (\_ -> let ?x = boom in rec1 3)+ strict ref "branchF" (\_ -> let ?x = boom in branch False 1)+ strict ref "local" (\_ -> let ?x = boom in local 1)+ strict ref "c3" (\_ -> let ?x = boom in c3 1)+ strict ref "closure" (\_ -> let ?x = boom in closure 1 2)+ strict ref "cfgUse" (\_ -> let ?cfg = error "boom" in cfgUse 1)+ strict ref "fnUse" (\_ -> let ?k = error "boom" in fnUse 1)++ failures <- readIORef ref+ if failures == 0+ then putStrLn "ALL OK"+ else do putStrLn (show failures ++ " failure(s)")+ exitWith (ExitFailure 1)++eq :: (Eq a, Show a) => IORef Int -> String -> a -> a -> IO ()+eq ref name got want = do+ r <- tryAny (evaluate (got == want))+ case r of+ Right True -> ok name+ Right False -> bad ref (name ++ ": got " ++ show got ++ ", want " ++ show want)+ Left e -> bad ref (name ++ ": exception " ++ show e)++-- | Check that forcing the expression raises the error in the implicit parameter.+strict :: Show a => IORef Int -> String -> (() -> a) -> IO ()+strict ref name k = do+ r <- tryErrorCall (evaluate (k ()))+ case r of+ Left _ -> ok (name ++ " forced")+ Right v -> bad ref (name ++ ": not forced, returned " ++ show v)++ok :: String -> IO ()+ok name = putStrLn ("ok " ++ name)++bad :: IORef Int -> String -> IO ()+bad ref msg = do+ putStrLn ("FAIL " ++ msg)+ modifyIORef' ref (+1)++tryAny :: IO a -> IO (Either SomeException a)+tryAny = try++tryErrorCall :: IO a -> IO (Either ErrorCall a)+tryErrorCall = try