packages feed

bloodhound-1.0.0.0: tests/Test/MlTrainedModelsSpec.hs

{-# LANGUAGE OverloadedStrings #-}

module Test.MlTrainedModelsSpec (spec) where

import Data.Aeson
import Data.Aeson.KeyMap qualified as KeyMap
import Data.ByteString.Lazy.Char8 qualified as LBS
import Data.Scientific (Scientific)
import Database.Bloodhound.ElasticSearch9.Requests qualified as RequestsES9
import Database.Bloodhound.ElasticSearch9.Types qualified as Types
import TestsUtils.Import

hasKey :: Key -> LBS.ByteString -> Bool
hasKey k bs = case decode bs of
  Just (Object obj) -> k `KeyMap.member` obj
  _ -> False

spec :: Spec
spec =
  describe "ML trained model APIs (/_ml/trained_models/*)" $ do
    describe "MlTrainedModelId JSON" $ do
      it "round-trips an id" $ do
        let i = Types.MlTrainedModelId "lang_ident_model_1"
        decode (encode i) `shouldBe` Just i

      it "unMlTrainedModelId extracts the underlying Text" $
        Types.unMlTrainedModelId (Types.MlTrainedModelId "x")
          `shouldBe` ("x" :: Text)

    describe "MlTrainedModelPutBody JSON" $ do
      it "encodes a compressed-definition body, omitting optional fields" $ do
        let cfg = object ["regression" .= object []]
            body =
              Types.defaultMlTrainedModelPutBody
                { Types.mltmpInferenceConfig = cfg,
                  Types.mltmpCompressedDefinition = Just "base64blob"
                }
        hasKey "description" (encode body) `shouldBe` False
        hasKey "definition" (encode body) `shouldBe` False
        hasKey "compressed_definition" (encode body) `shouldBe` True
        decode (encode body) `shouldBe` Just body

      it "decodes a documented create body" $ do
        let raw =
              LBS.pack
                "{\"compressed_definition\":\"Zm9v\",\"description\":\"d\",\
                \\"inference_config\":{\"regression\":{}}}"
        case decode raw :: Maybe Types.MlTrainedModelPutBody of
          Just b -> do
            Types.mltmpCompressedDefinition b `shouldBe` Just ("Zm9v" :: Text)
            Types.mltmpDescription b `shouldBe` Just ("d" :: Text)
          Nothing -> expectationFailure "failed to decode put body"

      it "defaults a missing inference_config to {}" $
        case decode (LBS.pack "{\"compressed_definition\":\"x\"}") :: Maybe Types.MlTrainedModelPutBody of
          Just b -> Types.mltmpInferenceConfig b `shouldBe` object []
          Nothing -> expectationFailure "failed to decode put body"

    describe "MlTrainedModelUpdateBody JSON" $ do
      it "encodes to {} when empty" $
        encode Types.defaultMlTrainedModelUpdateBody
          `shouldBe` "{}"

      it "round-trips a deprecation update" $ do
        let b =
              Types.defaultMlTrainedModelUpdateBody
                { Types.mltmubDeprecated = Just True
                }
        decode (encode b) `shouldBe` Just b

    describe "MlTrainedModelOptions params" $ do
      it "renders nothing by default" $
        Types.mlTrainedModelOptionsParams Types.defaultMlTrainedModelOptions
          `shouldBe` []

      it "renders deploy-related params in order" $ do
        let opts =
              Types.defaultMlTrainedModelOptions
                { Types.mltooWaitForCompletion = Just True,
                  Types.mltooPriority = Just "normal",
                  Types.mltooThreads = Just 4
                }
        Types.mlTrainedModelOptionsParams opts
          `shouldBe` [ ("wait_for_completion", Just "true"),
                       ("priority", Just "normal"),
                       ("threads", Just "4")
                     ]

    describe "MlTrainedModelConfig JSON" $ do
      it "decodes a documented config verbatim" $ do
        let raw =
              LBS.pack
                "{\"model_id\":\"m\",\"created_by\":\"ml\",\"version\":\"17.0.0\",\
                \\"description\":\"d\",\"deprecated\":false,\
                \\"inference_config\":{\"regression\":{}},\
                \\"model_size_bytes\":1024}"
        case decode raw :: Maybe Types.MlTrainedModelConfig of
          Just c -> do
            Types.mltmcModelId c `shouldBe` Just ("m" :: Text)
            Types.mltmcDeprecated c `shouldBe` Just False
            Types.mltmcInferenceConfig c `shouldSatisfy` isJust
            extras <- pure (Types.mltmcExtras c)
            extras `shouldSatisfy` isJust
            forM_ extras $ \km ->
              KeyMap.lookup "model_size_bytes" km `shouldSatisfy` isJust
          Nothing -> expectationFailure "failed to decode config"

    describe "MlTrainedModelsResponse JSON" $ do
      it "decodes a list envelope, defaulting count" $ do
        let raw =
              LBS.pack
                "{\"trained_model_configs\":[{\"model_id\":\"a\"},{\"model_id\":\"b\"}]}"
        case decode raw :: Maybe Types.MlTrainedModelsResponse of
          Just r -> do
            length (Types.mltrConfigs r) `shouldBe` 2
            Types.mltrCount r `shouldBe` Just 2
          Nothing -> expectationFailure "failed to decode response"

      it "decodes an explicit count" $ do
        let raw = LBS.pack "{\"count\":5,\"trained_model_configs\":[]}"
        case decode raw :: Maybe Types.MlTrainedModelsResponse of
          Just r -> Types.mltrCount r `shouldBe` Just 5
          Nothing -> expectationFailure "failed to decode response"

    describe "MlTrainedModelStatsResponse JSON" $ do
      it "decodes a stats envelope, defaulting count" $ do
        let raw =
              LBS.pack
                "{\"trained_model_stats\":[{\"model_id\":\"a\",\"model_size_bytes\":2048}]}"
        case decode raw :: Maybe Types.MlTrainedModelStatsResponse of
          Just r -> do
            length (Types.mltsrStats r) `shouldBe` 1
            Types.mltsrCount r `shouldBe` Just 1
            case Types.mltsrStats r of
              (s : _) -> do
                Types.mltmsModelId s `shouldBe` Just ("a" :: Text)
                let extras = Types.mltmsExtras s
                extras `shouldSatisfy` isJust
                forM_ extras $ \km ->
                  KeyMap.size km `shouldBe` 1
              [] -> expectationFailure "expected one stats entry"
          Nothing -> expectationFailure "failed to decode stats response"

    describe "MlTrainedModelDefinition JSON" $ do
      it "decodes a definition envelope" $ do
        let raw =
              LBS.pack
                "{\"model_id\":\"m\",\"total_definition_length\":2048,\
                \\"model_definition\":{\"trained_model\":{\"tree\":{}}}}"
        case decode raw :: Maybe Types.MlTrainedModelDefinition of
          Just d -> do
            Types.mltmdModelId d `shouldBe` Just ("m" :: Text)
            Types.mltmdTotalDefinitionLength d `shouldBe` Just (2048 :: Scientific)
            Types.mltmdDefinition d `shouldSatisfy` isJust
            Types.mltmdExtras d `shouldBe` Nothing
          Nothing -> expectationFailure "failed to decode definition"

    describe "MlTrainedModelInferResponse JSON" $ do
      it "round-trips an arbitrary prediction body" $ do
        let body = object ["results" .= [object ["predicted_value" .= (1.5 :: Double)]]]
            r = Types.MlTrainedModelInferResponse body
        decode (encode r) `shouldBe` Just r

    describe "endpoint shape" $ do
      it "PUTs /_ml/trained_models/<id> with a body" $ do
        let req =
              RequestsES9.putMlTrainedModel
                "mymodel"
                Types.defaultMlTrainedModelPutBody
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel"]
        bhRequestBody req `shouldSatisfy` isJust

      it "GETs /_ml/trained_models (all)" $ do
        let req = RequestsES9.getMlTrainedModels Nothing
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models"]

      it "GETs /_ml/trained_models/<id>" $ do
        let req = RequestsES9.getMlTrainedModels (Just "mymodel")
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel"]

      it "GETs /_ml/trained_models/<id>/definition" $ do
        let req = RequestsES9.getMlTrainedModelDefinition "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "definition"]

      it "POSTs /_ml/trained_models/<id>/_update with a body" $ do
        let req =
              RequestsES9.updateMlTrainedModel
                "mymodel"
                Types.defaultMlTrainedModelUpdateBody
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "_update"]
        bhRequestBody req `shouldSatisfy` isJust

      it "DELETEs /_ml/trained_models/<id>" $ do
        let req = RequestsES9.deleteMlTrainedModel "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel"]
        bhRequestBody req `shouldSatisfy` isNothing

      it "POSTs /_ml/trained_models/<id>/deployment/_deploy (empty body)" $ do
        let req = RequestsES9.deployMlTrainedModel "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "deployment", "_deploy"]
        bhRequestBody req `shouldSatisfy` isJust

      it "POSTs /_ml/trained_models/<id>/deployment/_undeploy" $ do
        let req = RequestsES9.undeployMlTrainedModel "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "deployment", "_undeploy"]

      it "POSTs /_ml/trained_models/<id>/deployment/_start" $ do
        let req = RequestsES9.startMlTrainedModelDeployment "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "deployment", "_start"]

      it "POSTs /_ml/trained_models/<id>/deployment/_stop" $ do
        let req = RequestsES9.stopMlTrainedModelDeployment "mymodel"
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "deployment", "_stop"]

      it "GETs /_ml/trained_models/<id>/_stats" $ do
        let req = RequestsES9.getMlTrainedModelStats (Just "mymodel")
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "_stats"]

      it "GETs /_ml/trained_models/_stats (all)" $ do
        let req = RequestsES9.getMlTrainedModelStats Nothing
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "_stats"]

      it "POSTs /_ml/trained_models/<id>/_infer with a body" $ do
        let req =
              RequestsES9.inferMlTrainedModel
                "mymodel"
                (Types.MlTrainedModelInferRequest (object ["docs" .= [object ["f" .= (1 :: Int)]]]))
        getRawEndpoint (bhRequestEndpoint req)
          `shouldBe` ["_ml", "trained_models", "mymodel", "_infer"]
        bhRequestBody req `shouldSatisfy` isJust