{-# LANGUAGE QuasiQuotes #-} {-# LANGUAGE OverloadedRecordDot #-} {-# LANGUAGE OverloadedLabels #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE MultilineStrings #-} {-# LANGUAGE OverloadedLists #-} {-# LANGUAGE ApplicativeDo #-} {-# LANGUAGE RecursiveDo #-} {- HLINT ignore "Use camelCase" -} module Gyehoek.CPS.Lower (lower, lowerProgram) where import Gyehoek.CPS.Syntax import Data.Vector.Strict (Vector) import Control.Lens hiding (op) import Numeric.Natural import qualified Data.Vector.Strict as V import Gyehoek.Wasm qualified as Wasm import Gyehoek.Wasm hiding (Expr) import Language.Sexp.Located qualified as SL import Control.Monad.Fix import qualified Gyehoek.Sexp import Data.Text qualified as T import Data.Foldable (fold) import Gyehoek.Sexp (encodeOrShow) import Gyehoek.Prelude data Env = MkEnv { vars :: Vector Name , kvars :: Vector Name } deriving (Show, Generic) type instance Index Env = Natural type instance IxValue Env = Name instance Ixed Env where ix i = #vars . ix (fromIntegral i) tonat :: Integral a => a -> Natural tonat = fromIntegral -- | @makeSmallFixnum@ emits an expression injecting the i32 on top -- of the stack into the SCM unitype. makeSmallFixnum :: Wasm.Expr makeSmallFixnum = [expr| (@gyehoek "construct small fixnum") (i32.const 1) i32.shl ref.i31 |] getArgRegister :: Natural -> SL.Sexp getArgRegister n = SL.Symbol [i|$arg#{n}|] -- | Given an expression @e@ leaving a @ref eq@ atop the stack, -- @pushArg rt n e@ sets the nth slot of the arg-passing array to the -- result of @e@. pushArg :: Natural -> Wasm.Expr -> Wasm.Expr pushArg n e = [expr| (@gyehoek begin pushArg) ##{e} (global.set #{reg}) (@gyehoek end pushArg) |] where reg = getArgRegister n -- | Pop the nth arg from the arg-passing array onto the stack. popArg :: Natural -> Wasm.Expr popArg n = [expr| (@gyehoek begin popArg) (global.get #{reg}) ref.as_non_null (@gyehoek end popArg) |] where reg = getArgRegister n lowerVal :: (HasCallStack, GenMod :> es) => Env -> Val -> Eff es Wasm.Expr lowerVal g (ValImm imm) = pure $ case imm of ImmInt n -> [expr| (i32.const #{n}) ##{makeSmallFixnum} |] ImmBool b -> [expr| (i32.const #{b'}) ref.i31 |] where b' :: Int = if b then 0b11 else 0b01 _ -> _ lowerVal g (ValVar x) = do pure [expr|(global.get #{l})|] where l = getArgRegister . fromIntegral . succ $ V.elemIndex x g.vars ^?! _Just lower' :: (GenMod :> es) => Env -> Exp -> Eff es Wasm.Expr lower' g (Halt [v]) = do arg <- pushArg 0 <$> lowerVal g v pure [expr| ##{arg} (return_call $halt (i32.const 1)) |] lower' g e@(ExpPrim p k) = ([expr|(@gyehoek :origin #{origin})|]<>) <$> case p of PrimAdd x y -> lowerBinOp "i32.add" g x y k PrimMul x y -> lowerBinOp "i32.mul" g x y k where origin = encodeOrShow @_ @Text e lower' g (ExpIf c t f) = do c' <- lowerVal g c t' <- lower' g t f' <- lower' g f pure [expr| ##{c'} (call $gh-truthy?) (if (then ##{t'}) (else ##{f'})) |] lower' g (ExpLetRec [(r,AbsKappa kap)] e) = do idx <- lowerKappa g kap let g' = g & #kvars <>~ [r] e' <- lower' g' e let origin = encodeOrShow @_ @Text e pure [expr| (@gyehoek :origin #{origin}) (@gyehoek "push cont" :idx #{idx}) (array.set $cont-stack-type (global.get $cont-stack) (global.get $cont-stack-top) (ref.func #{idx})) (global.set $cont-stack-top (i32.add (global.get $cont-stack-top) (i32.const 1))) ##{e'} |] lower' g (ExpLetRec [(r,AbsLambda lam)] e) = do idx <- lowerLambda g lam let g' = g & #vars <>~ [r] let n = succ $ length g.vars e' <- lower' g' e let reg = getArgRegister . fromIntegral $ n pure [expr| (i32.const 0) (ref.func #{idx}) (struct.new $closure) (global.set #{reg}) ##{e'} |] lower' g e@(ExpApply f xs ktail) = do let nargs = length xs f' <- lowerVal g f let l = succ $ V.elemIndex ktail g.kvars ^?! _Just args <- fold <$> itraverse (\i -> fmap (pushArg $ tonat i) . lowerVal g) xs let origin = encodeOrShow @_ @Text e pure [expr| (@gyehoek :origin #{origin}) (@gyehoek "load args") ##{args} (i32.const 1) ##{f'} (ref.cast (ref $closure)) (struct.get $closure $code) (return_call_ref $cont-type) (@gyehoek todo (f' ##{f'}) (ktail #{l})) |] lower' g e@(ExpContinue k xs) = do let nargs = length xs args <- fold <$> itraverse (\i -> fmap (pushArg $ tonat i) . lowerVal g) xs let origin = encodeOrShow @_ @Text e pure [expr| (@gyehoek :origin #{origin}) (@gyehoek "push args") ##{args} (@gyehoek "nargs") (i32.const #{nargs}) (@gyehoek "pop cont stack") (global.get $cont-stack-top) (i32.const #{l}) i32.sub (global.set $cont-stack-top) (global.get $cont-stack) (global.get $cont-stack-top) (array.get $cont-stack-type) ref.as_non_null (return_call_ref $cont-type) |] where l = succ $ V.elemIndex k g.kvars ^?! _Just lower' g e = error $ case Gyehoek.Sexp.encode e of Left _ -> show e Right x -> T.unpack x lowerKappa :: GenMod :> es => Env -> Kappa -> Eff es Idx lowerKappa g e@(MkKappa xs m) = do let g' = g & #vars <>~ V.fromList xs m' <- lower' g' m let origin = encodeOrShow @_ @Text e idx <- Wasm.defineFunction [wat| (func (param i32) (@gyehoek :origin #{origin}) (local (ref eq) (ref eq) (ref eq) (ref eq) (ref eq)) ##{m'}) |] Wasm.emit [wats|(elem declare funcref (ref.func #{idx}))|] pure idx lowerLambda :: GenMod :> es => Env -> Lambda -> Eff es Idx lowerLambda g e@(MkLambda xs ktail m) = do let g' = g & #vars .~ V.fromList xs & #kvars <>~ [ktail] m' <- lower' g' m let origin = encodeOrShow @_ @Text e idx <- Wasm.defineFunction [wat| (func (param i32) (@gyehoek :origin #{origin}) (local (ref eq) (ref eq) (ref eq) (ref eq) (ref eq)) ##{m'}) |] Wasm.emit [wats|(elem declare funcref (ref.func #{idx}))|] pure idx lowerBinOp :: (GenMod :> es) => Text -> Env -> Val -> Val -> Kappa -> Eff es Wasm.Expr lowerBinOp op g x y (MkKappa [r] e) = do let op' = SL.Symbol op let g' = g & #vars <>~ [r] let n = succ $ length (g ^. #vars) let reg = getArgRegister . fromIntegral $ n x' <- lowerVal g x y' <- lowerVal g y e' <- lower' g' e pure [expr| ##{x'} (i31.get_s (ref.cast (ref i31))) (i32.const 1) i32.shr_u ##{y'} (i31.get_s (ref.cast (ref i31))) (i32.const 1) i32.shr_u #{op'} ##{makeSmallFixnum} (global.set #{reg}) ##{e'} |] emitRuntime :: GenMod :> es => Eff es () emitRuntime = mfix \runtime -> do Wasm.defineFunctions [wats| (import "gyehoek" "write" (func $gh-write (param (ref eq)))) (import "gyehoek" "truthy?" (func $gh-truthy? (param (ref eq)) (result i32))) |] -- cont stack Wasm.defineTypes [wats| (type $heap-object (sub (struct (field $hash (mut i32))))) (type $cont-type (func (param i32))) (type $cont-stack-type (array (mut (ref null $cont-type)))) (type $closure (sub $heap-object (struct (field $hash (mut i32)) (field $code (ref $cont-type))))) |] Wasm.defineGlobals [wats| (global $cont-stack-top (mut i32) (i32.const 0)) (global $cont-stack (ref $cont-stack-type) (array.new_default $cont-stack-type (i32.const 128))) |] -- arg registers Wasm.defineGlobals [wats| (global $arg0 (mut (ref null eq)) (ref.null eq)) (global $arg1 (mut (ref null eq)) (ref.null eq)) (global $arg2 (mut (ref null eq)) (ref.null eq)) (global $arg3 (mut (ref null eq)) (ref.null eq)) (global $arg4 (mut (ref null eq)) (ref.null eq)) (global $arg5 (mut (ref null eq)) (ref.null eq)) (global $arg6 (mut (ref null eq)) (ref.null eq)) (global $arg7 (mut (ref null eq)) (ref.null eq)) (global $arg8 (mut (ref null eq)) (ref.null eq)) (global $arg9 (mut (ref null eq)) (ref.null eq)) (global $arg10 (mut (ref null eq)) (ref.null eq)) (global $arg11 (mut (ref null eq)) (ref.null eq)) (global $arg12 (mut (ref null eq)) (ref.null eq)) (global $arg13 (mut (ref null eq)) (ref.null eq)) (global $arg14 (mut (ref null eq)) (ref.null eq)) (global $arg15 (mut (ref null eq)) (ref.null eq)) |] -- other things 😼 Wasm.defineGlobal [wat| (global $result (mut (ref null eq)) (ref.null eq)) |] -- procedures let arg = popArg 0 Wasm.defineFunction [wat| (func $halt (param i32) ##{arg} (global.set $result)) |] pure () lower :: Exp -> Eff es Text lower e = fmap Wasm.renderModule . Wasm.execGenMod $ do runtime <- emitRuntime let g = MkEnv mempty mempty e' <- lower' g e let origin = encodeOrShow @_ @Text e Wasm.defineFunction [wat| (func $scm-entry (param i32) (@gyehoek :origin #{origin}) (local (ref eq) (ref eq) (ref eq) (ref eq) (ref eq)) ##{e'}) |] Wasm.defineFunction [wat| (func (export "main") (call $scm-entry (i32.const 0)) (call $gh-write (ref.as_non_null (global.get $result)))) |] lowerProgram :: Program -> Eff es Text lowerProgram (MkProgram e) = lower e