packages feed

langchain-hs-0.0.5.0: src/Langchain/Memory/Core.hs

{-# LANGUAGE DuplicateRecordFields #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE OverloadedStrings #-}

{- |
Module      : Langchain.Memory.Core
Description : Effect-polymorphic memory management for LangChain Haskell
Copyright   : (c) 2025-2026 Tushar Adhatrao
License     : MIT
Maintainer  : Tushar Adhatrao <tusharadhatrao@gmail.com>
Stability   : experimental

Thread-safe, effect-polymorphic conversation memory interfaces using STM.
-}
module Langchain.Memory.Core
  ( BaseMemory (..)
  , WindowBufferMemory (..)
  , newWindowBufferMemory
  , TokenBufferMemory (..)
  , newTokenBufferMemory
  , countTokens
  , trimMessages
  , initialMessages
  ) where

import Control.Concurrent.STM (TVar, atomically, modifyTVar', newTVarIO, readTVarIO, writeTVar)
import Control.Monad.Except (MonadError, throwError)
import Control.Monad.IO.Class (MonadIO, liftIO)
import Data.Text (Text)
import qualified Data.Text as T

import Langchain.Core.Error (LangchainError, memoryError)
import Langchain.Core.Model
  ( Message (..)
  , Role (..)
  , assistantMessage
  , extractMessageText
  , systemMessage
  , userMessage
  )

-- | Effect-polymorphic BaseMemory typeclass
class BaseMemory mem where
  -- | Retrieve current conversation messages
  messages ::
    (MonadIO m, MonadError LangchainError m) =>
    mem ->
    m [Message]

  -- | Add a user message to history
  addUserMessage ::
    (MonadIO m, MonadError LangchainError m) =>
    mem ->
    Text ->
    m ()
  addUserMessage mem txt = addMessage mem (userMessage txt)

  -- | Add an AI response message to history
  addAiMessage ::
    (MonadIO m, MonadError LangchainError m) =>
    mem ->
    Text ->
    m ()
  addAiMessage mem txt = addMessage mem (assistantMessage txt)

  -- | Add a structured message to history
  addMessage ::
    (MonadIO m, MonadError LangchainError m) =>
    mem ->
    Message ->
    m ()

  -- | Reset memory to initial state
  clear ::
    (MonadIO m, MonadError LangchainError m) =>
    mem ->
    m ()

-- | Sliding window memory backed by thread-safe STM TVar
data WindowBufferMemory = WindowBufferMemory
  { maxWindowSize :: !Int
  , memVar :: !(TVar [Message])
  }

instance Show WindowBufferMemory where
  show (WindowBufferMemory sz _) = "WindowBufferMemory { maxWindowSize = " ++ show sz ++ " }"

instance Eq WindowBufferMemory where
  (WindowBufferMemory sz1 tv1) == (WindowBufferMemory sz2 tv2) =
    sz1 == sz2 && tv1 == tv2

-- | Construct a thread-safe WindowBufferMemory in MonadIO
newWindowBufferMemory :: MonadIO m => Int -> [Message] -> m WindowBufferMemory
newWindowBufferMemory sz initMsgs = liftIO $ do
  tv <- newTVarIO initMsgs
  pure $ WindowBufferMemory sz tv

instance BaseMemory WindowBufferMemory where
  messages (WindowBufferMemory _ tv) = liftIO $ readTVarIO tv

  addMessage (WindowBufferMemory maxSz tv) newMsg = liftIO $ do
    atomically $ modifyTVar' tv $ \currMsgs ->
      let combined = currMsgs ++ [newMsg]
       in if length combined > maxSz
            then removeOldestNonSystem combined
            else combined
    where
      removeOldestNonSystem [] = []
      removeOldestNonSystem (m : ms)
        | messageRole m == System = m : removeOldestNonSystem ms
        | otherwise = ms

  clear (WindowBufferMemory _ tv) = liftIO $ do
    atomically $ writeTVar tv [systemMessage "You are a helpful AI assistant"]

-- | Token-based sliding window memory type
data TokenBufferMemory = TokenBufferMemory
  { maxTokens :: !Int
  , memVar :: !(TVar [Message])
  }

instance Show TokenBufferMemory where
  show (TokenBufferMemory maxT _) = "TokenBufferMemory { maxTokens = " ++ show maxT ++ " }"

instance Eq TokenBufferMemory where
  (TokenBufferMemory t1 tv1) == (TokenBufferMemory t2 tv2) =
    t1 == t2 && tv1 == tv2

-- | Construct a new TokenBufferMemory
newTokenBufferMemory :: MonadIO m => Int -> [Message] -> m TokenBufferMemory
newTokenBufferMemory maxT initMsgs = liftIO $ do
  tv <- newTVarIO initMsgs
  pure $ TokenBufferMemory maxT tv

-- | Approximate token count: 4 characters ≈ 1 token
countTokens :: [Message] -> Int
countTokens = sum . map (\m -> ceiling (fromIntegral (T.length (extractMessageText m)) / (4.0 :: Double)))

instance BaseMemory TokenBufferMemory where
  messages (TokenBufferMemory _ tv) = liftIO $ readTVarIO tv

  addMessage (TokenBufferMemory maxT tv) newMsg = do
    let newMsgTokens = countTokens [newMsg]
    if newMsgTokens > maxT
      then
        throwError $
          memoryError "New message exceeds maximum token limit" (Just "TokenBufferMemory") Nothing
      else liftIO $ atomically $ modifyTVar' tv $ \currMsgs ->
        trimToLimit currMsgs newMsgTokens [newMsg]
    where
      trimToLimit currMsgs newMsgToks acc =
        let candidate = currMsgs ++ acc
         in if countTokens candidate <= maxT
              then candidate
              else case removeOldestNonSystem currMsgs of
                Just trimmed -> trimToLimit trimmed newMsgToks acc
                Nothing -> [newMsg]

      removeOldestNonSystem [] = Nothing
      removeOldestNonSystem (m : ms)
        | messageRole m == System = fmap (m :) (removeOldestNonSystem ms)
        | otherwise = Just ms

  clear (TokenBufferMemory _ tv) = liftIO $ do
    atomically $ writeTVar tv [systemMessage "You are a helpful AI assistant"]

-- | Pure helper to trim messages to last N
trimMessages :: Int -> [Message] -> [Message]
trimMessages n msgs = drop (max 0 (length msgs - n)) msgs

-- | Pure helper to construct initial system message history
initialMessages :: Text -> [Message]
initialMessages sysPrompt = [systemMessage sysPrompt]