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/HomeAssistant/Runtime/Flags.hs b/src/HomeAssistant/Runtime/Flags.hs new file mode 100644 index 0000000..876a3a9 --- /dev/null +++ b/src/HomeAssistant/Runtime/Flags.hs @@ -0,0 +1,89 @@ +{-# 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.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 + putStrLn ("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 + putStrLn ("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/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