diff --git a/default.nix b/default.nix index ed421bf..672ab13 100644 --- a/default.nix +++ b/default.nix @@ -1,8 +1,8 @@ { mkDerivation, aeson, annotated-exception, async, base, bytestring , cereal, cereal-conduit, conduit, containers, directory, ekg-core -, exceptions, filepath, hedgehog, hspec, hspec-hedgehog, http-types -, katip, lens, lens-aeson, lib, network, process, retry, scotty -, servant, servant-server, stm, text, time, unliftio +, exceptions, filepath, hedgehog, hspec, hspec-hedgehog, http-media +, http-types, katip, lens, lens-aeson, lib, network, process, retry +, servant, servant-server, stm, temporary, text, time, unliftio , unordered-containers, uuid, wai, warp, websockets }: mkDerivation { @@ -14,15 +14,15 @@ mkDerivation { libraryHaskellDepends = [ aeson annotated-exception async base bytestring cereal cereal-conduit conduit containers directory ekg-core exceptions - filepath http-types katip lens lens-aeson network process retry - scotty servant servant-server stm text time unliftio + filepath http-media http-types katip lens lens-aeson network + process retry servant servant-server stm text time unliftio unordered-containers uuid wai warp websockets ]; executableHaskellDepends = [ base ]; testHaskellDepends = [ aeson annotated-exception async base bytestring cereal containers - directory ekg-core hedgehog hspec hspec-hedgehog katip process stm - text time unordered-containers uuid + directory ekg-core hedgehog hspec hspec-hedgehog katip process + servant stm temporary text time unordered-containers uuid ]; license = lib.meta.getLicenseFromSpdxId "BSD-3-Clause"; mainProgram = "home-assistant-controller"; diff --git a/home-assistant-controller.cabal b/home-assistant-controller.cabal index a655a1b..d0e333e 100644 --- a/home-assistant-controller.cabal +++ b/home-assistant-controller.cabal @@ -70,6 +70,7 @@ library , HomeAssistant.Runtime , HomeAssistant.Runtime.Bus , HomeAssistant.Runtime.Connection + , HomeAssistant.Runtime.Flags , HomeAssistant.Runtime.Graphing , HomeAssistant.Runtime.Metrics , HomeAssistant.Runtime.Supervisor @@ -161,6 +162,7 @@ test-suite home-assistant-controller-test , BedroomSpec , BusSpec , ConnectionSpec + , FlagsSpec , GraphingSpec , MetricsSpec , RuntimeSpec @@ -200,4 +202,5 @@ test-suite home-assistant-controller-test ekg-core, unordered-containers, process, - directory + directory, + temporary diff --git a/src/AFRP.hs b/src/AFRP.hs index e3855e7..3a0bbba 100644 --- a/src/AFRP.hs +++ b/src/AFRP.hs @@ -103,9 +103,10 @@ mergeCodec (Codec agetter aputter) (Codec bgetter bputter) = Codec (mergeGet age data Pair a b = Pair !a !b data Request = Request - { requestTime :: !UTCTime - , requestTimeZone :: !TimeZone - , requestTraceId :: !UUID + { requestTime :: !UTCTime + , requestTimeZone :: !TimeZone + , requestTraceId :: !UUID + , requestHandler :: !T.Text } deriving (Show, Eq) diff --git a/src/HomeAssistant/Runtime.hs b/src/HomeAssistant/Runtime.hs index 891ec51..c8a191c 100644 --- a/src/HomeAssistant/Runtime.hs +++ b/src/HomeAssistant/Runtime.hs @@ -21,6 +21,7 @@ import Data.Time (getCurrentTime, getCurrentTimeZone) import Data.Void (Void) import HomeAssistant.Controller (HASS, HASSEff (..)) import HomeAssistant.Runtime.Bus +import HomeAssistant.Runtime.Flags (loadFlags) import HomeAssistant.Runtime.Connection (writerAction, readerAction) import Network.Socket (withSocketsDo) import System.Environment (getEnv, lookupEnv) @@ -43,34 +44,34 @@ import HomeAssistant.Runtime.Supervisor (supervised) import HomeAssistant.Controller.Hallway (hallwayLightsController) import qualified HttpServer -step :: (MonadIO m) => FilePath -> UUID -> Auto m a b -> a -> m (b, Auto m a b) -step path trace st a = do +step :: (MonadIO m) => FilePath -> T.Text -> UUID -> Auto m a b -> a -> m (b, Auto m a b) +step path name trace st a = do now <- liftIO getCurrentTime tz <- liftIO getCurrentTimeZone - let req = Request now tz trace + let req = Request now tz trace name stepAutoSerializing path st req a -data Controller = forall b. Controller Bool T.Text (HASS (Event Value) b) +data Controller = forall b. Controller T.Text (HASS (Event Value) b) controllers :: [Controller] controllers = - [ Controller True "bedroom-presence" bedroomPresenceController - , Controller True "bedroom-button" bedroomButtonController - , Controller True "bedroom-drawer" bedroomDrawerController - , Controller True "bedroom-humidifier" humidifierController - , Controller True "school-light-controller" schoolLightController - , Controller True "kitchen-motion-controller" kitchenMotionController - , Controller True "livingroom-presence" livingroomPresenceController - , Controller True "hallway-motion-controller" hallwayLightsController - , Controller True "children-button" childrenBedroomButtonController - , Controller True "livingroom-plants" livingroomPlantLights + [ Controller "bedroom-presence" bedroomPresenceController + , Controller "bedroom-button" bedroomButtonController + , Controller "bedroom-drawer" bedroomDrawerController + , Controller "bedroom-humidifier" humidifierController + , Controller "school-light-controller" schoolLightController + , Controller "kitchen-motion-controller" kitchenMotionController + , Controller "livingroom-presence" livingroomPresenceController + , Controller "hallway-motion-controller" hallwayLightsController + , Controller "children-button" childrenBedroomButtonController + , Controller "livingroom-plants" livingroomPlantLights ] -- | Steps the machine for every inbound message; service calls go to the -- bus. A restart re-dups the inbound channel and starts from the machine's -- initial state; messages broadcast during the restart window are lost. runController :: MonadIO m => FilePath -> Bus -> Controller -> m Void -runController rootDir bus (Controller _enabled name machine ) = liftIO $ do +runController rootDir bus (Controller name machine ) = liftIO $ do inbound <- atomically (dupTChan (busInbound bus)) let ns = Namespace [name] let workerDefinition = runMealy machine (runKatipContextT (busLogEnv bus) () ns . channelHassEval bus) @@ -83,7 +84,7 @@ runController rootDir bus (Controller _enabled name machine ) = liftIO $ do go path inbound f = do msg <- atomically (readTChan inbound) uuid <- UUID.V4.nextRandom - (_, next) <- step path uuid f msg + (_, next) <- step path name uuid f msg go path inbound next lookupPort :: IO Int @@ -102,22 +103,23 @@ defaultMain = withSocketsDo $ do System.Metrics.registerGcMetrics store appMetrics <- HomeAssistant.Runtime.Metrics.registerAppMetrics store rateLimitMetrics <- registerRateLimitMetrics store - withBus severity appMetrics $ \bus -> do + rootPath <- fromMaybe "/tmp/" <$> lookupEnv "HA_LIB_DIR" + let allNames = map (\(Controller n _) -> n) controllers + flags <- loadFlags rootPath allNames + withBus severity appMetrics flags $ \bus -> do token <- getEnv "HA_TOKEN" host <- getEnv "HA_HOST" - rootPath <- fromMaybe "/tmp/" <$> lookupEnv "HA_LIB_DIR" - rrdPath <- fromMaybe "hass-controller.rrd" <$> lookupEnv "HA_RRD_PATH" + let rrdPath = rootPath "hass-controller.rrd" rrdtool <- fromMaybe "rrdtool" <$> lookupEnv "HA_RRDTOOL" metricsPort <- lookupPort writerLimiter <- slidingWindowLimiter rateLimitMetrics 10 20 - let active = [c | c@(Controller True _ _) <- controllers] - ents = foldMap (\(Controller _ _ m) -> entities m) active + let ents = foldMap (\(Controller _ m) -> entities m) controllers workers = [ ("reader", readerAction host 8123 token ents bus) , ("writer", writerAction writerLimiter bus) , ("metrics", HomeAssistant.Runtime.Metrics.metricsAction store rrdPath rrdtool) - , ("metrics-http", HttpServer.runHttpServer (busLogEnv bus) rrdPath rrdtool metricsPort) - ] ++ [ (name, runController rootPath bus c) | c@(Controller True name _) <- active ] + , ("metrics-http", HttpServer.runHttpServer (busLogEnv bus) rrdPath rrdtool (busFlags bus) metricsPort) + ] ++ [ (name, runController rootPath bus c) | c@(Controller name _) <- controllers ] runKatipContextT (busLogEnv bus) () mempty $ mapConcurrently_ (uncurry supervised) workers diff --git a/src/HomeAssistant/Runtime/Bus.hs b/src/HomeAssistant/Runtime/Bus.hs index 23509f1..74bcb61 100644 --- a/src/HomeAssistant/Runtime/Bus.hs +++ b/src/HomeAssistant/Runtime/Bus.hs @@ -31,7 +31,9 @@ import System.IO (stdout) import System.Metrics.Counter (inc) import Data.UUID (toText) import AFRP (Request(..), Event(..)) +import Control.Monad (when) import Control.Monad.IO.Class (MonadIO, liftIO) +import HomeAssistant.Runtime.Flags (Flags, isEnabled) -- | Shared runtime state: inbound is a broadcast channel (controllers -- read from 'dupTChan' copies), outbound queues service calls for the @@ -43,10 +45,11 @@ data Bus = Bus , busGen :: CallIdGen , busLogEnv :: LogEnv , busMetrics :: AppMetrics + , busFlags :: Flags } -withBus :: Severity -> AppMetrics -> (Bus -> IO a) -> IO a -withBus severity metrics callback = do +withBus :: Severity -> AppMetrics -> Flags -> (Bus -> IO a) -> IO a +withBus severity metrics flags callback = do handleScribe <- mkHandleScribe ColorIfTerminal stdout (permitItem severity) V2 let makeLogEnv = registerScribe "stdout" handleScribe defaultScribeSettings =<< initLogEnv "hass-controller" "production" -- closeScribes will stop accepting new logs, flush existing ones and clean up resources @@ -58,6 +61,7 @@ withBus severity metrics callback = do <*> mkCallIdGen 0 <*> pure le <*> pure metrics + <*> pure flags callback bus recordInbound :: Bus -> IO () @@ -70,7 +74,8 @@ channelHassEval :: (MonadIO m, KatipContext m) => Bus -> HASSEff a -> m a channelHassEval bus = \case CallService req svc -> katipAddContext (sl "traceId" (toText (requestTraceId req))) $ do logFM DebugS (ls $ show svc) - liftIO $ atomically $ writeTChan (busOutbound bus) (req, svc) + enabled <- isEnabled (busFlags bus) (requestHandler req) + when enabled $ liftIO $ atomically $ writeTChan (busOutbound bus) (req, svc) Debug x -> logFM DebugS (ls $ show x) Trace req x -> katipAddContext (sl "traceId" (toText (requestTraceId req))) $ logFM InfoS (ls $ show x) diff --git a/src/HomeAssistant/Runtime/Flags.hs b/src/HomeAssistant/Runtime/Flags.hs new file mode 100644 index 0000000..c594e9e --- /dev/null +++ b/src/HomeAssistant/Runtime/Flags.hs @@ -0,0 +1,90 @@ +{-# LANGUAGE OverloadedStrings #-} + +module HomeAssistant.Runtime.Flags + ( Name + , Flags(..) + , loadFlags + , isEnabled + , setEnabled + , allFlags + , validNames + ) where + +import Control.Concurrent.MVar (MVar, newMVar, withMVar) +import Control.Concurrent.STM (TVar, atomically, newTVarIO, readTVarIO, stateTVar) +import Control.Exception (IOException, handle) +import Control.Monad.IO.Class (MonadIO, liftIO) +import qualified Conduit as C +import Conduit ((.|)) +import Data.Aeson (decodeFileStrict', encode) +import qualified Data.ByteString.Lazy as BL +import Data.Map.Strict (Map) +import qualified Data.Map.Strict as M +import Data.Set (Set) +import qualified Data.Set as S +import Data.Text (Text) +import System.FilePath (()) +import System.IO (hPutStrLn, stderr) +import System.IO.Error (isDoesNotExistError) + +type Name = Text + +data Flags = Flags + { flagsState :: TVar (Map Name Bool) + , flagsLock :: MVar () + , flagsPath :: FilePath + , flagsValidNames :: Set Text + } + +loadFlags :: MonadIO m => FilePath -> [Name] -> m Flags +loadFlags rootPath allNames = liftIO $ do + let path = rootPath "handlers.json" + fileMap <- readFlagsFile path + let defaults = M.fromList [(n, True) | n <- allNames] + initial = M.unionWith (\fileVal _default -> fileVal) fileMap defaults + tv <- newTVarIO initial + lock <- newMVar () + pure Flags + { flagsState = tv + , flagsLock = lock + , flagsPath = path + , flagsValidNames = S.fromList allNames + } + +readFlagsFile :: FilePath -> IO (Map Name Bool) +readFlagsFile path = + handle onMissingOrUnreadable $ do + decoded <- decodeFileStrict' path + case decoded of + Nothing -> do + hPutStrLn stderr ("handlers.json is unparseable; treating all handlers as enabled: " <> path) + pure mempty + Just m -> pure (maybe mempty id m) + where + onMissingOrUnreadable :: IOException -> IO (Map Name Bool) + onMissingOrUnreadable e + | isDoesNotExistError e = pure mempty + | otherwise = do + hPutStrLn stderr ("could not read handlers.json (" <> show e <> "); treating all handlers as enabled") + pure mempty + +isEnabled :: MonadIO m => Flags -> Name -> m Bool +isEnabled flags name = liftIO $ + M.findWithDefault True name <$> readTVarIO (flagsState flags) + +allFlags :: MonadIO m => Flags -> m (Map Name Bool) +allFlags flags = liftIO $ readTVarIO (flagsState flags) + +setEnabled :: MonadIO m => Flags -> Name -> Bool -> m () +setEnabled Flags{flagsState, flagsLock, flagsPath} name val = + liftIO $ withMVar flagsLock $ \_ -> do + m <- atomically $ stateTVar flagsState $ \cur -> + let next = M.insert name val cur in (next, next) + safeWrite flagsPath m + +validNames :: Flags -> Set Text +validNames = flagsValidNames + +safeWrite :: FilePath -> Map Name Bool -> IO () +safeWrite path m = C.runResourceT $ C.runConduit $ + C.yieldMany [BL.toStrict (encode m)] .| C.sinkFileCautious path diff --git a/src/HttpServer.hs b/src/HttpServer.hs index 6a9258d..6d4f0a5 100644 --- a/src/HttpServer.hs +++ b/src/HttpServer.hs @@ -7,13 +7,17 @@ import Control.Monad.IO.Class (liftIO, MonadIO) import qualified Data.ByteString as BS import Data.Void (Void) import Network.Wai.Handler.Warp (run) -import Servant.API ((:>), Capture, Get, QueryParam, (:-), MimeRender (..), Accept (..), NamedRoutes) +import Servant.API ((:>), Capture, Get, QueryParam, (:-), MimeRender (..), Accept (..), NamedRoutes, JSON, ReqBody, Put) import Data.ByteString (ByteString) import GHC.Generics (Generic) import Servant (Handler, err404) -import Control.Monad.Catch (throwM, MonadCatch, MonadThrow) +import Control.Monad.Catch (throwM, catch, MonadCatch, MonadThrow, SomeException) import Servant.Server (err500, ServerT) import Data.Maybe (fromMaybe) +import Data.Map.Strict (Map) +import qualified Data.Set as S +import Data.Text (Text) +import HomeAssistant.Runtime.Flags (Flags, allFlags, setEnabled, validNames) import Control.Exception.Annotated (checkpoint, Annotation (Annotation)) import Servant.Server.Generic (genericServeT) import Network.Wai (Application) @@ -30,12 +34,27 @@ instance Accept ImagePng where contentType _ = "image" // "png" -newtype API mode = API +data API mode = API { getMetrics :: mode :- "metrics" :> Capture "graph" Graph :> QueryParam "range" Range :> Get '[ImagePng] ByteString + , getFlags :: mode :- "flags" :> Get '[JSON] (Map Text Bool) + , putFlag :: mode :- "flags" :> Capture "name" Text :> ReqBody '[JSON] Bool :> Put '[JSON] (Map Text Bool) } deriving Generic +getFlagsHandler :: Flags -> LoggingHandler (Map Text Bool) +getFlagsHandler flags = allFlags flags + +putFlagHandler :: Flags -> Text -> Bool -> LoggingHandler (Map Text Bool) +putFlagHandler flags name val + | name `S.notMember` validNames flags = throwM err404 + | otherwise = do + setEnabled flags name val `catch` \e -> do + logFM ErrorS (ls ("failed to persist handler flag: " <> show (e :: SomeException))) + throwM err500 + allFlags flags + + getMetricsHandler :: FilePath -> FilePath -> Graph -> Maybe Range -> LoggingHandler ByteString getMetricsHandler rrdPath rrdtool graphMode mRange = checkpoint (Annotation (graphMode, mRange)) $ do case lookupGraph graphs graphMode of @@ -46,14 +65,18 @@ getMetricsHandler rrdPath rrdtool graphMode mRange = checkpoint (Annotation (gra either (\e -> logFM ErrorS (ls e) >> throwM err500) pure g -- server :: FilePath -> FilePath -> API AsServer -server :: FilePath -> FilePath -> ServerT (NamedRoutes API) LoggingHandler -server rrdPath rrdtool = API { getMetrics = getMetricsHandler rrdPath rrdtool } +server :: FilePath -> FilePath -> Flags -> ServerT (NamedRoutes API) LoggingHandler +server rrdPath rrdtool flags = API + { getMetrics = getMetricsHandler rrdPath rrdtool + , getFlags = getFlagsHandler flags + , putFlag = putFlagHandler flags + } -servantApp :: LogEnv -> FilePath -> FilePath -> Application -servantApp le rrdPath rrdtool = genericServeT (toHandler le) (server rrdPath rrdtool) +servantApp :: LogEnv -> FilePath -> FilePath -> Flags -> Application +servantApp le rrdPath rrdtool flags = genericServeT (toHandler le) (server rrdPath rrdtool flags) -runHttpServer :: MonadIO m => LogEnv -> FilePath -> FilePath -> Int -> m Void -runHttpServer le rrdPath rrdtool port = liftIO $ forever $ run port (servantApp le rrdPath rrdtool) +runHttpServer :: MonadIO m => LogEnv -> FilePath -> FilePath -> Flags -> Int -> m Void +runHttpServer le rrdPath rrdtool flags port = liftIO $ forever $ run port (servantApp le rrdPath rrdtool flags) newtype LoggingHandler a = LoggingHandler (KatipContextT Handler a) deriving (Functor, Applicative, Monad, MonadIO, Katip, KatipContext, MonadCatch, MonadThrow) via (KatipContextT Handler) diff --git a/test/AFRPSpec.hs b/test/AFRPSpec.hs index c47c463..f201dc3 100644 --- a/test/AFRPSpec.hs +++ b/test/AFRPSpec.hs @@ -21,7 +21,7 @@ import Test.Hspec import Test.Hspec.Hedgehog fakeRequest :: Request -fakeRequest = Request (sec 0) utc nil +fakeRequest = Request (sec 0) utc nil "test" sec :: Integer -> UTCTime sec n = UTCTime (toEnum 0) (fromIntegral n) @@ -42,7 +42,7 @@ runTimed m = go (runMealy m id) where go _ [] = [] go w ((s, a) : as) = - case runIdentity (stepAuto w (Request (sec s) utc nil) a) of + case runIdentity (stepAuto w (Request (sec s) utc nil "test") a) of (b, w') -> b : go w' as -- | A minimal State monad for observing effectful arrows (e.g. whenA gating). diff --git a/test/BusSpec.hs b/test/BusSpec.hs index a8c37d8..616a86f 100644 --- a/test/BusSpec.hs +++ b/test/BusSpec.hs @@ -6,6 +6,7 @@ import AFRP (Request (..), Event (..)) import Control.Concurrent.STM ( atomically , dupTChan + , isEmptyTChan , readTChan , writeTChan ) @@ -15,9 +16,11 @@ import Data.Time (UTCTime (..), utc) import Data.UUID (nil) import HomeAssistant.Controller (HASSEff (..), Service (..), Target(..)) import HomeAssistant.Runtime.Bus +import HomeAssistant.Runtime.Flags (loadFlags, setEnabled) import HomeAssistant.Runtime.Metrics (registerAppMetrics) import Katip (Namespace (Namespace), runKatipContextT, Severity (..)) import qualified System.Metrics as Metrics +import System.IO.Temp (withTempDirectory) import Test.Hspec spec :: Spec @@ -33,13 +36,22 @@ spec = describe "Bus" $ do r2 `shouldBe` (Event (Number 1), Event (Number 2)) it "channelHassEval writes CallService to the outbound channel" $ withTestBus $ \bus -> do - let req = Request (UTCTime (toEnum 0) 0) utc nil + let req = Request (UTCTime (toEnum 0) 0) utc nil "test" svc = Service "light" "turn_on" Nothing [EntityId "light.bedroom_masse"] runKatipContextT (busLogEnv bus) () (Namespace ["test"]) $ channelHassEval bus (CallService req svc) (_, svc') <- atomically $ readTChan (busOutbound bus) svc' `shouldBe` svc + it "channelHassEval drops CallService when the handler is disabled" $ + withTestFlagsBus $ \bus -> do + setEnabled (busFlags bus) "test-handler" False + let req = Request (UTCTime (toEnum 0) 0) utc nil "test-handler" + svc = Service "light" "turn_on" Nothing [EntityId "light.test"] + runKatipContextT (busLogEnv bus) () (Namespace ["test"]) $ + channelHassEval bus (CallService req svc) + atomically (isEmptyTChan (busOutbound bus)) `shouldReturn` True + it "generates unique sequential call ids" $ do gen <- mkCallIdGen 0 @@ -51,16 +63,28 @@ spec = describe "Bus" $ do it "increments the trigger and service counters" $ do store <- Metrics.newStore m <- registerAppMetrics store - withBus InfoS m $ \bus -> do - recordInbound bus - recordInbound bus - recordOutbound bus - sample <- Metrics.sampleAll store - HM.lookup "hass.trigger.in" sample `shouldBe` Just (Metrics.Counter 2) - HM.lookup "hass.service.out" sample `shouldBe` Just (Metrics.Counter 1) + withTempDirectory "/tmp" "bus-spec" $ \dir -> do + flags <- loadFlags dir [] + withBus InfoS m flags $ \bus -> do + recordInbound bus + recordInbound bus + recordOutbound bus + sample <- Metrics.sampleAll store + HM.lookup "hass.trigger.in" sample `shouldBe` Just (Metrics.Counter 2) + HM.lookup "hass.service.out" sample `shouldBe` Just (Metrics.Counter 1) withTestBus :: (Bus -> IO a) -> IO a withTestBus action = do store <- Metrics.newStore m <- registerAppMetrics store - withBus InfoS m action + withTempDirectory "/tmp" "bus-spec" $ \dir -> do + flags <- loadFlags dir [] + withBus InfoS m flags action + +withTestFlagsBus :: (Bus -> IO a) -> IO a +withTestFlagsBus action = do + store <- Metrics.newStore + m <- registerAppMetrics store + withTempDirectory "/tmp" "bus-spec" $ \dir -> do + flags <- loadFlags dir ["test-handler"] + withBus InfoS m flags action diff --git a/test/ConnectionSpec.hs b/test/ConnectionSpec.hs index 733fb36..f4c30a0 100644 --- a/test/ConnectionSpec.hs +++ b/test/ConnectionSpec.hs @@ -47,7 +47,7 @@ spec = do req :: Int -> Request req n = Request (UTCTime (toEnum 0) (fromIntegral (0 :: Int))) utc - (fromJust (fromString uuid)) + (fromJust (fromString uuid)) "test" where pad i = replicate (12 - length (show i)) '0' <> show i uuid = "00000000-0000-0000-0000-" <> pad n diff --git a/test/FlagsSpec.hs b/test/FlagsSpec.hs new file mode 100644 index 0000000..73fa0fe --- /dev/null +++ b/test/FlagsSpec.hs @@ -0,0 +1,75 @@ +{-# LANGUAGE OverloadedStrings #-} + +module FlagsSpec (spec) where + +import Control.Concurrent.Async (mapConcurrently_) +import Control.Monad.IO.Class (liftIO) +import Data.Aeson (decodeFileStrict) +import Data.Map.Strict (Map) +import qualified Data.Map.Strict as M +import Data.Text (Text) +import qualified Hedgehog.Gen as Gen +import qualified Hedgehog.Range as Range +import System.IO.Temp (withTempDirectory) +import Test.Hspec +import Test.Hspec.Hedgehog + +import HomeAssistant.Runtime.Flags + +spec :: Spec +spec = describe "Flags" $ do + describe "loadFlags" $ do + it "defaults all names to True when the file is missing (opt-out)" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + flags <- loadFlags dir ["a", "b"] + m <- allFlags flags + m `shouldBe` M.fromList [("a" :: Text, True), ("b", True)] + + it "lets the file override the True default, missing keys stay True" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + let path = dir <> "/handlers.json" + writeFile path "{\"b\": false}" + flags <- loadFlags dir ["a", "b"] + m <- allFlags flags + m `shouldBe` M.fromList [("a" :: Text, True), ("b" :: Text, False)] + + it "treats a corrupt file as empty (all-enabled) without throwing" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + let path = dir <> "/handlers.json" + writeFile path "not json at all" + flags <- loadFlags dir ["a"] + m <- allFlags flags + m `shouldBe` M.fromList [("a" :: Text, True)] + + describe "isEnabled" $ do + it "returns the stored value" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + let path = dir <> "/handlers.json" + writeFile path "{\"a\": false}" + flags <- loadFlags dir ["a"] + isEnabled flags "a" `shouldReturn` False + + it "defaults unknown names to True" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + flags <- loadFlags dir [] + isEnabled flags "unknown" `shouldReturn` True + + describe "setEnabled" $ do + it "updates the in-memory map and persists to disk" $ + withTempDirectory "/tmp" "flags-spec" $ \dir -> do + flags <- loadFlags dir ["a"] + setEnabled flags "a" False + isEnabled flags "a" `shouldReturn` False + -- reload from disk to confirm persistence + flags' <- loadFlags dir ["a"] + isEnabled flags' "a" `shouldReturn` False + + it "concurrent writes leave the file consistent with the TVar" $ + hedgehog $ do + n <- forAll $ Gen.int (Range.linear 1 50) + liftIO $ withTempDirectory "/tmp" "flags-spec" $ \dir -> do + flags <- loadFlags dir ["x"] + mapConcurrently_ (\i -> setEnabled flags "x" (even i)) [1 .. n] + m <- allFlags flags + decoded <- decodeFileStrict (dir <> "/handlers.json") + (decoded :: Maybe (Map Text Bool)) `shouldBe` Just m diff --git a/test/Main.hs b/test/Main.hs index fdd0264..a783121 100644 --- a/test/Main.hs +++ b/test/Main.hs @@ -6,6 +6,7 @@ import qualified AFRPSpec import qualified BedroomSpec import qualified BusSpec import qualified ConnectionSpec +import qualified FlagsSpec import qualified GraphingSpec import qualified MetricsSpec import qualified RuntimeSpec @@ -17,6 +18,7 @@ main = hspec $ do BedroomSpec.spec BusSpec.spec ConnectionSpec.spec + FlagsSpec.spec GraphingSpec.spec MetricsSpec.spec RuntimeSpec.spec diff --git a/test/Support.hs b/test/Support.hs index ff09e2b..781652c 100644 --- a/test/Support.hs +++ b/test/Support.hs @@ -20,7 +20,7 @@ import AFRP (Mealy (..), Request (..), stepAuto) import HomeAssistant.Controller (HASSEff (..), Service) fakeRequest :: Request -fakeRequest = Request (sec 0) utc nil +fakeRequest = Request (sec 0) utc nil "test" sec :: Integer -> UTCTime sec n = UTCTime (toEnum 0) (fromIntegral n)