from __future__ import annotations
import uuid
from typing import Any, Dict, List, Optional, Sequence
from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.messages import (
AIMessage,
BaseMessage,
HumanMessage,
SystemMessage,
messages_from_dict,
messages_to_dict,
)
_ROLE_HUMAN = "human"
_ROLE_AI = "ai"
_ROLE_SYSTEM = "system"
_LANGCHAIN_ROLE_MAP = {
"human": _ROLE_HUMAN,
"ai": _ROLE_AI,
"system": _ROLE_SYSTEM,
}
_AGENTDB_TO_LC_ROLE: Dict[str, str] = {v: k for k, v in _LANGCHAIN_ROLE_MAP.items()}
def _message_to_role(message: BaseMessage) -> str:
if isinstance(message, HumanMessage):
return _ROLE_HUMAN
if isinstance(message, AIMessage):
return _ROLE_AI
if isinstance(message, SystemMessage):
return _ROLE_SYSTEM
return message.type
def _row_to_message(row: Dict[str, Any]) -> BaseMessage:
role = row.get("role", "")
content = row.get("content", "")
if role == _ROLE_HUMAN:
return HumanMessage(content=content)
if role == _ROLE_AI:
return AIMessage(content=content)
if role == _ROLE_SYSTEM:
return SystemMessage(content=content)
return HumanMessage(content=content)
class AgentDBChatMessageHistory(BaseChatMessageHistory):
def __init__(
self,
db_path: str,
conversation_id: Optional[str] = None,
title: Optional[str] = None,
) -> None:
try:
import agentdb as _agentdb
except ImportError as exc: raise ImportError(
"The 'datacules-agentdb' package is required. "
"Install it with: pip install datacules-agentdb"
) from exc
self._db_path = db_path
self._conversation_id = conversation_id or str(uuid.uuid4())
self._title = title
self._db = _agentdb.AgentDB.open(db_path)
try:
self._db.create_conversation(
self._conversation_id,
title=self._title,
)
except RuntimeError:
pass
@property
def messages(self) -> List[BaseMessage]:
rows = self._db.get_messages(self._conversation_id)
return [_row_to_message(row) for row in rows]
def add_message(self, message: BaseMessage) -> None:
role = _message_to_role(message)
content = message.content if isinstance(message.content, str) else str(message.content)
self._db.add_message(
self._conversation_id,
role,
content,
)
def add_messages(self, messages: Sequence[BaseMessage]) -> None:
for message in messages:
self.add_message(message)
def clear(self) -> None:
self._db.delete_conversation(self._conversation_id)
self._db.create_conversation(
self._conversation_id,
title=self._title,
)
def add_user_message(self, message: str) -> None:
self.add_message(HumanMessage(content=message))
def add_ai_message(self, message: str) -> None:
self.add_message(AIMessage(content=message))
@property
def conversation_id(self) -> str:
return self._conversation_id
def __len__(self) -> int:
return len(self.messages)
def __repr__(self) -> str:
return (
f"AgentDBChatMessageHistory("
f"db_path={self._db_path!r}, "
f"conversation_id={self._conversation_id!r})"
)
try:
from langchain_core.memory import BaseMemory
class AgentDBChatMemory(BaseMemory):
db_path: str
conversation_id: str = ""
memory_key: str = "history"
input_key: str = "input"
output_key: str = "output"
return_messages: bool = False
class Config:
arbitrary_types_allowed = True
def __init__(self, **data: Any) -> None:
if not data.get("conversation_id"):
data["conversation_id"] = str(uuid.uuid4())
super().__init__(**data)
object.__setattr__(
self,
"_history",
AgentDBChatMessageHistory(
db_path=self.db_path,
conversation_id=self.conversation_id,
),
)
@property
def memory_variables(self) -> List[str]:
return [self.memory_key]
def load_memory_variables(self, inputs: Dict[str, Any]) -> Dict[str, Any]:
messages = self._history.messages if self.return_messages:
return {self.memory_key: messages}
lines = []
for msg in messages:
prefix = "Human" if isinstance(msg, HumanMessage) else "AI"
lines.append(f"{prefix}: {msg.content}")
return {self.memory_key: "\n".join(lines)}
def save_context(
self, inputs: Dict[str, Any], outputs: Dict[str, Any]
) -> None:
human_text = inputs.get(self.input_key, "")
ai_text = outputs.get(self.output_key, "")
if human_text:
self._history.add_user_message(str(human_text)) if ai_text:
self._history.add_ai_message(str(ai_text))
def clear(self) -> None:
self._history.clear()
except ImportError:
AgentDBChatMemory = None