packages feed

llm-simple-0.2.0.0: test/LLM/ChatSpec.hs

{-# OPTIONS_GHC -Wno-missing-fields #-}

module LLM.ChatSpec (spec) where

import Data.Aeson (object, (.=))
import Data.IORef (IORef, modifyIORef', newIORef, readIORef)
import Data.Map qualified as Map
import Data.Text (Text)
import Heptapod (generate)
import LLM.Agent.Events (noEventObserver)
import LLM.Agent.Generate (generateText, streamText)
import LLM.Agent.Types
  ( Agent (..),
    RuntimeArgs (..),
    Tool (..),
    ToolMap,
  )
import LLM.Core.Abort (AbortSignal, abort, newAbortSignal)
import LLM.Core.Types
  ( CacheHint (..),
    ChatRequest (..),
    ChatResponse (..),
    ContentPart (..),
    PartBody (..),
    textPart, toolCallPart, pattern UserTurn,
    Turn (..),
    cacheEphemeral,
    imageUrlPart,
    LLMError (..),
    LLMGateway (..),
    LLMHooks (..),
    ToolDef (ToolDef, toolDescription, toolName, toolParameters, toolReadonly),
    mkToolCall,
  )
import LLM.Core.Usage (PricingInfo (..), defaultPricingInfo, mkUsage)
import LLM.Generate.Logger (noHooks)
import LLM.Generate.ModelConfig
  ( ModelCapabilities (..),
    ModelConfig (..),
    ModelWithFallbacks (..),
    defaultModelCapabilities,
  )
import LLM.Generate.Types
  ( GenerateError (..),
    GenerateErrorResult (..),
    GenerateTextResult (..),
  )
import Test.Hspec
  ( Spec,
    describe,
    expectationFailure,
    it,
    shouldBe,
    shouldSatisfy,
  )

-- | A mock gateway that returns a fixed response
mockGateway :: ChatResponse -> LLMGateway
mockGateway resp =
  LLMGateway
    { gwName = "mock",
      gwGenerateText = \_ _ -> pure (Right resp),
      gwStreamText = \_ _ _ -> pure (Right resp),
      gwGenerateObject = \_ _ _ -> pure (Right (object [], Nothing))
    }

-- | A mock gateway that returns an error
mockErrorGateway :: LLMError -> LLMGateway
mockErrorGateway err =
  LLMGateway
    { gwName = "mock-error",
      gwGenerateText = \_ _ -> pure (Left err),
      gwStreamText = \_ _ _ -> pure (Left err),
      gwGenerateObject = \_ _ _ -> pure (Left err)
    }

-- | A mock gateway that calls a tool, then responds with text
mockToolGateway :: LLMGateway
mockToolGateway =
  LLMGateway
    { gwName = "mock-tool",
      gwGenerateText = \_ req ->
        if any isToolTurn req.reqConversation
          then pure $ Right (ChatResponse "The weather is sunny." [textPart "The weather is sunny."] (Just (mkUsage 80 15)) Nothing)
          else
            let tc = mkToolCall "call_1" "get_weather" (object ["location" .= ("London" :: Text)])
             in pure $ Right (ChatResponse "" [toolCallPart tc] (Just (mkUsage 50 10)) Nothing),
      gwStreamText = \_ _ _ -> pure $ Right (ChatResponse "" [] Nothing Nothing),
      gwGenerateObject = \_ _ _ -> pure $ Right (object [], Nothing)
    }
  where
    isToolTurn (ToolTurn _) = True
    isToolTurn _ = False

zeroPricing :: PricingInfo
zeroPricing = defaultPricingInfo 0 0

-- | Wrap a gateway in a ModelConfig with test defaults
mockModel :: LLMGateway -> ModelConfig
mockModel gw =
  ModelConfig
    { mcGateway = gw,
      mcModel = "test-model",
      mcPricing = zeroPricing,
      mcMaxTokens = 1024,
      mcTemperature = Nothing,
      mcThinking = Nothing,
      mcCapabilities = defaultModelCapabilities,
      mcRequestTimeout = Nothing,
      mcThrottleDelay = Nothing,
      mcRetryCount = 0,
      mcJitterBackoff = 0
    }

weatherTool :: Tool Text
weatherTool =
  Tool
    { toolDef =
        ToolDef
          { toolName = "get_weather",
            toolDescription = "Get weather",
            toolReadonly = True,
            toolParameters = object ["type" .= ("object" :: Text)]
          },
      toolExecute = \_ _ -> pure "Sunny, 22°C"
    }

defaultAgent :: Agent
defaultAgent =
  Agent
    { agName = "test",
      agSystemPrompt = Nothing,
      agTools = [],
      agMaxToolRounds = 10,
      agContextWindow = Nothing
    }

mkRuntime :: Maybe AbortSignal -> IO RuntimeArgs
mkRuntime mSig = do
  uuid <- generate
  pure
    RuntimeArgs
      { rtGenerationId = uuid,
        rtAbortSignal = mSig,
        rtLLMHooks =
          LLMHooks
            { onLLMRequest = \_ _ -> pure (),
              onLLMResponse = \_ _ -> pure (),
              onLLMResponseError = \_ _ -> pure ()
            },
        rtHooks = noHooks,
        rtOnEvent = noEventObserver,
        rtReadonly = False
      }

runGenerate ::
  Agent ->
  ModelWithFallbacks ->
  ToolMap Text ->
  Maybe AbortSignal ->
  [Turn] ->
  IO (Either GenerateErrorResult GenerateTextResult)
runGenerate agent models toolMap mSig turns = do
  rt <- mkRuntime mSig
  generateText agent models toolMap rt turns

runStreamGenerate ::
  Agent ->
  ModelWithFallbacks ->
  ToolMap Text ->
  Maybe AbortSignal ->
  [Turn] ->
  IO (Either GenerateErrorResult GenerateTextResult)
runStreamGenerate agent models toolMap mSig turns = do
  rt <- mkRuntime mSig
  streamText (\_ -> pure ()) agent models toolMap rt turns

-- | Gateway that records each request conversation, then behaves like 'mockToolGateway'.
capturingToolGateway :: IORef [[Turn]] -> LLMGateway
capturingToolGateway ref =
  let respond req = do
        modifyIORef' ref (req.reqConversation :)
        if any isToolTurn req.reqConversation
          then pure $ Right (ChatResponse "The weather is sunny." [textPart "The weather is sunny."] (Just (mkUsage 80 15)) Nothing)
          else
            let tc = mkToolCall "call_1" "get_weather" (object ["location" .= ("London" :: Text)])
             in pure $ Right (ChatResponse "" [toolCallPart tc] (Just (mkUsage 50 10)) Nothing)
   in LLMGateway
        { gwName = "mock-tool-capture",
          gwGenerateText = \_ -> respond,
          gwStreamText = \_ req _onEvent -> respond req,
          gwGenerateObject = \_ _ _ -> pure $ Right (object [], Nothing)
        }
  where
    isToolTurn (ToolTurn _) = True
    isToolTurn _ = False

hasCachedUserPrefix :: [Turn] -> Bool
hasCachedUserPrefix =
  any
    ( \case
        UserMessage parts ->
          any
            ( \case
                ContentPart (TextPart _) (Just CacheEphemeral) -> True
                _ -> False
            )
            parts
        _ -> False
    )

spec :: Spec
spec = describe "Chat" $ do
  let toolMap = Map.fromList [("get_weather", weatherTool)]
  describe "generateText" $ do
    it "returns text for a simple response" $ do
      let gw = mockGateway (ChatResponse "Hi there!" [textPart "Hi there!"] (Just (mkUsage 10 5)) Nothing)
          models = ModelWithFallbacks (mockModel gw) []
      result <- runGenerate defaultAgent models toolMap Nothing [UserTurn "hello"]
      case result of
        Right r -> do
          r.gtrText `shouldBe` "Hi there!"
          length r.gtrNewMessages `shouldBe` 1 -- assistantTurn
          r.gtrUsage `shouldBe` mkUsage 10 5
        Left err -> expectationFailure $ show err

    it "propagates errors" $ do
      let gw = mockErrorGateway (HttpError 500 "internal error")
          models = ModelWithFallbacks (mockModel gw) []
      result <- runGenerate defaultAgent models toolMap Nothing [UserTurn "hello"]
      case result of
        Left GenerateErrorResult {gerError = GErrLLM (HttpError 500 _)} -> pure ()
        other -> expectationFailure $ "Expected HttpError 500, got: " <> show other

    it "handles tool call loop" $ do
      let agent = defaultAgent {agTools = ["get_weather"]}
          models = ModelWithFallbacks (mockModel mockToolGateway) []
      result <- runGenerate agent models toolMap Nothing [UserTurn "weather in london?"]
      case result of
        Right r -> do
          r.gtrText `shouldBe` "The weather is sunny."
          -- assistantTurn(tool call) + ToolTurn + assistantTurn(final)
          length r.gtrNewMessages `shouldBe` 3
          r.gtrUsage `shouldBe` mkUsage 130 25 -- 50+80 input, 10+15 output
        Left err -> expectationFailure $ show err

    it "preserves cache hints on user parts across non-streaming tool rounds" $ do
      ref <- newIORef []
      let agent = defaultAgent {agTools = ["get_weather"]}
          models = ModelWithFallbacks (mockModel (capturingToolGateway ref)) []
          msgs =
            [ UserMessage
                [ cacheEphemeral (textPart "static system-like context"),
                  textPart "weather in london?"
                ]
            ]
      result <- runGenerate agent models toolMap Nothing msgs
      case result of
        Right r -> do
          r.gtrText `shouldBe` "The weather is sunny."
          convs <- reverse <$> readIORef ref
          length convs `shouldBe` 2
          mapM_ (`shouldSatisfy` hasCachedUserPrefix) convs
        Left err -> expectationFailure $ show err

    it "respects maxToolRounds" $ do
      let infiniteToolGateway =
            LLMGateway
              { gwName = "mock-infinite",
                gwGenerateText = \_ _ ->
                  let tc = mkToolCall "call_1" "get_weather" (object [])
                   in pure $ Right (ChatResponse "" [toolCallPart tc] Nothing Nothing),
                gwStreamText = \_ _ _ -> pure $ Right (ChatResponse "" [] Nothing Nothing),
                gwGenerateObject = \_ _ _ -> pure $ Right (object [], Nothing)
              }
          agent = defaultAgent {agMaxToolRounds = 2, agTools = ["get_weather"]}
          models = ModelWithFallbacks (mockModel infiniteToolGateway) []
      result <- runGenerate agent models toolMap Nothing [UserTurn "test"]
      case result of
        Left GenerateErrorResult {gerError = GErrToolExceeded} -> pure ()
        other -> expectationFailure $ "Expected GErrToolExceeded, got: " <> show other

    it "falls back to next model on retryable error" $ do
      let failGw = mockErrorGateway (HttpError 503 "service unavailable")
          okGw = mockGateway (ChatResponse "Fallback worked!" [textPart "Fallback worked!"] (Just (mkUsage 10 5)) Nothing)
          models = ModelWithFallbacks (mockModel failGw) [mockModel okGw]
      result <- runGenerate defaultAgent models toolMap Nothing [UserTurn "hello"]
      case result of
        Right r -> r.gtrText `shouldBe` "Fallback worked!"
        Left err -> expectationFailure $ "Expected fallback success, got: " <> show err

    it "falls back on non-retryable error too" $ do
      let failGw = mockErrorGateway (HttpError 400 "bad request")
          okGw = mockGateway (ChatResponse "Fallback worked!" [textPart "Fallback worked!"] (Just (mkUsage 10 5)) Nothing)
          models = ModelWithFallbacks (mockModel failGw) [mockModel okGw]
      result <- runGenerate defaultAgent models toolMap Nothing [UserTurn "hello"]
      case result of
        Right r -> r.gtrText `shouldBe` "Fallback worked!"
        Left err -> expectationFailure $ "Expected fallback success, got: " <> show err

    it "skips a non-vision model and succeeds with a vision-capable fallback" $ do
      let boomGw =
            LLMGateway
              { gwName = "should-not-run",
                gwGenerateText = \_ _ -> pure (Left (HttpError 500 "should not be called")),
                gwStreamText = \_ _ _ -> pure (Left (HttpError 500 "should not be called")),
                gwGenerateObject = \_ _ _ -> pure (Left (HttpError 500 "should not be called"))
              }
          okGw = mockGateway (ChatResponse "I see a cat." [textPart "I see a cat."] (Just (mkUsage 10 5)) Nothing)
          noVision = mockModel boomGw
          withVision =
            (mockModel okGw)
              { mcCapabilities = defaultModelCapabilities {capVision = True}
              }
          models = ModelWithFallbacks noVision [withVision]
          msgs = [UserMessage [imageUrlPart "https://example.com/cat.png", textPart "what is this?"]]
      result <- runGenerate defaultAgent models toolMap Nothing msgs
      case result of
        Right r -> r.gtrText `shouldBe` "I see a cat."
        Left err -> expectationFailure $ "Expected vision fallback success, got: " <> show err

    it "returns UnsupportedCapability when no vision-capable model remains" $ do
      let noVision1 = mockModel (mockErrorGateway (HttpError 500 "unused"))
          noVision2 = mockModel (mockErrorGateway (HttpError 500 "unused"))
          models = ModelWithFallbacks noVision1 [noVision2]
          msgs = [UserMessage [imageUrlPart "https://example.com/cat.png", textPart "what?"]]
      result <- runGenerate defaultAgent models toolMap Nothing msgs
      case result of
        Left GenerateErrorResult {gerError = GErrLLM (UnsupportedCapability _)} -> pure ()
        other -> expectationFailure $ "Expected UnsupportedCapability, got: " <> show other

    it "returns error from last model when all fail" $ do
      let failGw1 = mockErrorGateway (HttpError 503 "service unavailable")
          failGw2 = mockErrorGateway (HttpError 400 "bad request")
          models = ModelWithFallbacks (mockModel failGw1) [mockModel failGw2]
      result <- runGenerate defaultAgent models toolMap Nothing [UserTurn "hello"]
      case result of
        Left GenerateErrorResult {gerError = GErrLLM (HttpError 400 _)} -> pure ()
        other -> expectationFailure $ "Expected HttpError 400 from last model, got: " <> show other

    it "returns Aborted when signal is fired before the call" $ do
      let gw = mockGateway (ChatResponse "Hi!" [textPart "Hi!"] Nothing Nothing)
          models = ModelWithFallbacks (mockModel gw) []
      sig <- newAbortSignal
      abort sig
      result <- runGenerate defaultAgent models toolMap (Just sig) [UserTurn "hello"]
      case result of
        Left GenerateErrorResult {gerError = GErrAborted} -> pure ()
        other -> expectationFailure $ "Expected GErrAborted, got: " <> show other

    it "returns Aborted during tool execution" $ do
      sig <- newAbortSignal
      let slowTool =
            Tool
              { toolDef =
                  ToolDef
                    { toolName = "slow",
                      toolDescription = "A slow tool",
                      toolReadonly = True,
                      toolParameters = object ["type" .= ("object" :: Text)]
                    },
                toolExecute = \_ _ -> do
                  abort sig
                  pure "done"
              }
          twoCallGw =
            LLMGateway
              { gwName = "mock-two",
                gwGenerateText = \_ _ ->
                  let tc1 = mkToolCall "c1" "slow" (object [])
                      tc2 = mkToolCall "c2" "slow" (object [])
                   in pure $ Right (ChatResponse "" [toolCallPart tc1, toolCallPart tc2] Nothing Nothing),
                gwStreamText = \_ _ _ -> pure $ Right (ChatResponse "" [] Nothing Nothing),
                gwGenerateObject = \_ _ _ -> pure $ Right (object [], Nothing)
              }
          tm = Map.fromList [("slow", slowTool)]
          agent = defaultAgent {agTools = ["slow"]}
          models = ModelWithFallbacks (mockModel twoCallGw) []
      result <- runGenerate agent models tm (Just sig) [UserTurn "go"]
      case result of
        Left GenerateErrorResult {gerError = GErrAborted} -> pure ()
        other -> expectationFailure $ "Expected GErrAborted during tools, got: " <> show other

    it "does not fall back on Aborted" $ do
      let gw = mockGateway (ChatResponse "Hi!" [textPart "Hi!"] Nothing Nothing)
          okGw = mockGateway (ChatResponse "Fallback" [textPart "Fallback"] Nothing Nothing)
          models = ModelWithFallbacks (mockModel gw) [mockModel okGw]
      sig <- newAbortSignal
      abort sig
      result <- runGenerate defaultAgent models Map.empty (Just sig) [UserTurn "hello"]
      case result of
        Left GenerateErrorResult {gerError = GErrAborted} -> pure ()
        other -> expectationFailure $ "Expected GErrAborted (no fallback), got: " <> show other

  describe "streamText" $ do
    it "preserves cache hints on user parts across streaming tool rounds" $ do
      ref <- newIORef []
      let agent = defaultAgent {agTools = ["get_weather"]}
          models = ModelWithFallbacks (mockModel (capturingToolGateway ref)) []
          msgs =
            [ UserMessage
                [ cacheEphemeral (textPart "static system-like context"),
                  textPart "weather in london?"
                ]
            ]
      result <- runStreamGenerate agent models toolMap Nothing msgs
      case result of
        Right r -> do
          r.gtrText `shouldBe` "The weather is sunny."
          convs <- reverse <$> readIORef ref
          length convs `shouldBe` 2
          mapM_ (`shouldSatisfy` hasCachedUserPrefix) convs
        Left err -> expectationFailure $ show err