codekingpro/portable-devtools
114k
1import json2from typing import List3 4from langchain_core.chat_history import BaseChatMessageHistory5from langchain_core.messages import (6 BaseMessage,7 message_to_dict,8 messages_from_dict,9)10 11 12class XataChatMessageHistory(BaseChatMessageHistory):13 """Chat message history stored in a Xata database."""14 15 def __init__(16 self,17 session_id: str,18 db_url: str,19 api_key: str,20 branch_name: str = "main",21 table_name: str = "messages",22 create_table: bool = True,23 ) -> None:24 """Initialize with Xata client."""25 try:26 from xata.client import XataClient27 except ImportError:28 raise ImportError(29 "Could not import xata python package. "30 "Please install it with `pip install xata`."31 )32 self._client = XataClient(33 api_key=api_key, db_url=db_url, branch_name=branch_name34 )35 self._table_name = table_name36 self._session_id = session_id37 38 if create_table:39 self._create_table_if_not_exists()40 41 def _create_table_if_not_exists(self) -> None:42 r = self._client.table().get_schema(self._table_name)43 if r.status_code <= 299:44 return45 if r.status_code != 404:46 raise Exception(47 f"Error checking if table exists in Xata: {r.status_code} {r}"48 )49 r = self._client.table().create(self._table_name)50 if r.status_code > 299:51 raise Exception(f"Error creating table in Xata: {r.status_code} {r}")52 r = self._client.table().set_schema(53 self._table_name,54 payload={55 "columns": [56 {"name": "sessionId", "type": "string"},57 {"name": "type", "type": "string"},58 {"name": "role", "type": "string"},59 {"name": "content", "type": "text"},60 {"name": "name", "type": "string"},61 {"name": "additionalKwargs", "type": "json"},62 ]63 },64 )65 if r.status_code > 299:66 raise Exception(f"Error setting table schema in Xata: {r.status_code} {r}")67 68 def add_message(self, message: BaseMessage) -> None:69 """Append the message to the Xata table"""70 msg = message_to_dict(message)71 r = self._client.records().insert(72 self._table_name,73 {74 "sessionId": self._session_id,75 "type": msg["type"],76 "content": message.content,77 "additionalKwargs": json.dumps(message.additional_kwargs),78 "role": msg["data"].get("role"),79 "name": msg["data"].get("name"),80 },81 )82 if r.status_code > 299:83 raise Exception(f"Error adding message to Xata: {r.status_code} {r}")84 85 @property86 def messages(self) -> List[BaseMessage]: # type: ignore[override]87 r = self._client.data().query(88 self._table_name,89 payload={90 "filter": {91 "sessionId": self._session_id,92 },93 "sort": {"xata.createdAt": "asc"},94 },95 )96 if r.status_code != 200:97 raise Exception(f"Error running query: {r.status_code} {r}")98 msgs = messages_from_dict(99 [100 {101 "type": m["type"],102 "data": {103 "content": m["content"],104 "role": m.get("role"),105 "name": m.get("name"),106 "additional_kwargs": json.loads(m["additionalKwargs"]),107 },108 }109 for m in r["records"]110 ]111 )112 return msgs113 114 def clear(self) -> None:115 """Delete session from Xata table."""116 while True:117 r = self._client.data().query(118 self._table_name,119 payload={120 "columns": ["id"],121 "filter": {122 "sessionId": self._session_id,123 },124 },125 )126 if r.status_code != 200:127 raise Exception(f"Error running query: {r.status_code} {r}")128 ids = [rec["id"] for rec in r["records"]]129 if len(ids) == 0:130 break131 operations = [132 {"delete": {"table": self._table_name, "id": id}} for id in ids133 ]134 self._client.records().transaction(payload={"operations": operations})135 