Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
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 
codekingpro/portable-devtools · Team Ai