chore: import upstream snapshot with attribution
CI / Run CI (push) Has been cancelled
CI / check-backend (push) Has been cancelled
CI / check-frontend (push) Has been cancelled
CI / tests (push) Has been cancelled
CI / e2e-tests (push) Has been cancelled
Copilot Setup Steps / copilot-setup-steps (push) Has been cancelled
CI / Run CI (push) Has been cancelled
CI / check-backend (push) Has been cancelled
CI / check-frontend (push) Has been cancelled
CI / tests (push) Has been cancelled
CI / e2e-tests (push) Has been cancelled
Copilot Setup Steps / copilot-setup-steps (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,111 @@
|
||||
import os
|
||||
import warnings
|
||||
from typing import Optional
|
||||
|
||||
from .base import BaseDataLayer
|
||||
from .utils import (
|
||||
queue_until_user_message as queue_until_user_message, # TODO: Consider deprecating re-export.; Redundant alias tells type checkers to STFU.
|
||||
)
|
||||
|
||||
_data_layer: Optional[BaseDataLayer] = None
|
||||
_data_layer_initialized = False
|
||||
|
||||
|
||||
def get_data_layer():
|
||||
global _data_layer, _data_layer_initialized
|
||||
|
||||
if not _data_layer_initialized:
|
||||
if _data_layer:
|
||||
# Data layer manually set, warn user that this is deprecated.
|
||||
|
||||
warnings.warn(
|
||||
"Setting data layer manually is deprecated. Use @data_layer instead.",
|
||||
DeprecationWarning,
|
||||
)
|
||||
|
||||
else:
|
||||
from chainlit.config import config
|
||||
|
||||
if config.code.data_layer:
|
||||
# When @data_layer is configured, call it to get data layer.
|
||||
_data_layer = config.code.data_layer()
|
||||
elif database_url := os.environ.get("DATABASE_URL"):
|
||||
from .chainlit_data_layer import ChainlitDataLayer
|
||||
|
||||
if os.environ.get("LITERAL_API_KEY"):
|
||||
warnings.warn(
|
||||
"Both LITERAL_API_KEY and DATABASE_URL specified. Ignoring Literal AI data layer and relying on data layer pointing to DATABASE_URL."
|
||||
)
|
||||
|
||||
bucket_name = os.environ.get("BUCKET_NAME")
|
||||
|
||||
# AWS S3
|
||||
aws_region = os.getenv("APP_AWS_REGION")
|
||||
aws_access_key = os.getenv("APP_AWS_ACCESS_KEY")
|
||||
aws_secret_key = os.getenv("APP_AWS_SECRET_KEY")
|
||||
dev_aws_endpoint = os.getenv("DEV_AWS_ENDPOINT")
|
||||
is_using_s3 = bool(aws_access_key and aws_secret_key and aws_region)
|
||||
|
||||
# Google Cloud Storage
|
||||
gcs_project_id = os.getenv("APP_GCS_PROJECT_ID")
|
||||
gcs_client_email = os.getenv("APP_GCS_CLIENT_EMAIL")
|
||||
gcs_private_key = os.getenv("APP_GCS_PRIVATE_KEY")
|
||||
is_using_gcs = bool(gcs_project_id)
|
||||
|
||||
# Azure Storage
|
||||
azure_storage_account = os.getenv("APP_AZURE_STORAGE_ACCOUNT")
|
||||
azure_storage_key = os.getenv("APP_AZURE_STORAGE_ACCESS_KEY")
|
||||
is_using_azure = bool(azure_storage_account and azure_storage_key)
|
||||
|
||||
storage_client = None
|
||||
|
||||
if sum([is_using_s3, is_using_gcs, is_using_azure]) > 1:
|
||||
warnings.warn(
|
||||
"Multiple storage configurations detected. Please use only one."
|
||||
)
|
||||
elif is_using_s3:
|
||||
from chainlit.data.storage_clients.s3 import S3StorageClient
|
||||
|
||||
storage_client = S3StorageClient(
|
||||
bucket=bucket_name,
|
||||
region_name=aws_region,
|
||||
aws_access_key_id=aws_access_key,
|
||||
aws_secret_access_key=aws_secret_key,
|
||||
endpoint_url=dev_aws_endpoint,
|
||||
)
|
||||
elif is_using_gcs:
|
||||
from chainlit.data.storage_clients.gcs import GCSStorageClient
|
||||
|
||||
storage_client = GCSStorageClient(
|
||||
project_id=gcs_project_id,
|
||||
client_email=gcs_client_email,
|
||||
private_key=gcs_private_key,
|
||||
bucket_name=bucket_name,
|
||||
)
|
||||
elif is_using_azure:
|
||||
from chainlit.data.storage_clients.azure_blob import (
|
||||
AzureBlobStorageClient,
|
||||
)
|
||||
|
||||
storage_client = AzureBlobStorageClient(
|
||||
container_name=bucket_name,
|
||||
storage_account=azure_storage_account,
|
||||
storage_key=azure_storage_key,
|
||||
)
|
||||
|
||||
_data_layer = ChainlitDataLayer(
|
||||
database_url=database_url, storage_client=storage_client
|
||||
)
|
||||
elif api_key := os.environ.get("LITERAL_API_KEY"):
|
||||
# When LITERAL_API_KEY is defined, use Literal AI data layer
|
||||
from .literalai import LiteralDataLayer
|
||||
|
||||
# support legacy LITERAL_SERVER variable as fallback
|
||||
server = os.environ.get("LITERAL_API_URL") or os.environ.get(
|
||||
"LITERAL_SERVER"
|
||||
)
|
||||
_data_layer = LiteralDataLayer(api_key=api_key, server=server)
|
||||
|
||||
_data_layer_initialized = True
|
||||
|
||||
return _data_layer
|
||||
@@ -0,0 +1,19 @@
|
||||
from fastapi import HTTPException
|
||||
|
||||
from chainlit.data import get_data_layer
|
||||
|
||||
|
||||
async def is_thread_author(username: str, thread_id: str):
|
||||
data_layer = get_data_layer()
|
||||
if not data_layer:
|
||||
raise HTTPException(status_code=400, detail="Data layer not initialized")
|
||||
|
||||
thread_author = await data_layer.get_thread_author(thread_id)
|
||||
|
||||
if not thread_author:
|
||||
raise HTTPException(status_code=404, detail="Thread not found")
|
||||
|
||||
if thread_author != username:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
else:
|
||||
return True
|
||||
@@ -0,0 +1,124 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional
|
||||
|
||||
from chainlit.types import (
|
||||
Feedback,
|
||||
PaginatedResponse,
|
||||
Pagination,
|
||||
ThreadDict,
|
||||
ThreadFilter,
|
||||
)
|
||||
|
||||
from .utils import queue_until_user_message
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from chainlit.element import Element, ElementDict
|
||||
from chainlit.step import StepDict
|
||||
from chainlit.user import PersistedUser, User
|
||||
|
||||
|
||||
class BaseDataLayer(ABC):
|
||||
"""Base class for data persistence."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_user(self, identifier: str) -> Optional["PersistedUser"]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def create_user(self, user: "User") -> Optional["PersistedUser"]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def delete_feedback(
|
||||
self,
|
||||
feedback_id: str,
|
||||
) -> bool:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_feedback(
|
||||
self,
|
||||
feedback: Feedback,
|
||||
) -> str:
|
||||
pass
|
||||
|
||||
@queue_until_user_message()
|
||||
@abstractmethod
|
||||
async def create_element(self, element: "Element"):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_element(
|
||||
self, thread_id: str, element_id: str
|
||||
) -> Optional["ElementDict"]:
|
||||
pass
|
||||
|
||||
@queue_until_user_message()
|
||||
@abstractmethod
|
||||
async def delete_element(self, element_id: str, thread_id: Optional[str] = None):
|
||||
pass
|
||||
|
||||
@queue_until_user_message()
|
||||
@abstractmethod
|
||||
async def create_step(self, step_dict: "StepDict"):
|
||||
pass
|
||||
|
||||
@queue_until_user_message()
|
||||
@abstractmethod
|
||||
async def update_step(self, step_dict: "StepDict"):
|
||||
pass
|
||||
|
||||
@queue_until_user_message()
|
||||
@abstractmethod
|
||||
async def delete_step(self, step_id: str):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_thread_author(self, thread_id: str) -> str:
|
||||
return ""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_thread(self, thread_id: str):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def list_threads(
|
||||
self, pagination: "Pagination", filters: "ThreadFilter"
|
||||
) -> "PaginatedResponse[ThreadDict]":
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_thread(self, thread_id: str) -> "Optional[ThreadDict]":
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def update_thread(
|
||||
self,
|
||||
thread_id: str,
|
||||
name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
metadata: Optional[Dict] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def build_debug_url(self) -> str:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def close(self) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_favorite_steps(self, user_id: str) -> List["StepDict"]:
|
||||
pass
|
||||
|
||||
async def set_step_favorite(
|
||||
self, step_dict: "StepDict", favorite: bool
|
||||
) -> "StepDict":
|
||||
metadata = step_dict.get("metadata") or {}
|
||||
metadata["favorite"] = favorite
|
||||
step_dict["metadata"] = metadata
|
||||
await self.update_step(step_dict)
|
||||
return step_dict
|
||||
@@ -0,0 +1,740 @@
|
||||
import json
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
import aiofiles
|
||||
import asyncpg # type: ignore
|
||||
|
||||
from chainlit.data.base import BaseDataLayer
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient
|
||||
from chainlit.data.utils import queue_until_user_message
|
||||
from chainlit.element import ElementDict
|
||||
from chainlit.logger import logger
|
||||
from chainlit.step import StepDict
|
||||
from chainlit.types import (
|
||||
Feedback,
|
||||
FeedbackDict,
|
||||
PageInfo,
|
||||
PaginatedResponse,
|
||||
Pagination,
|
||||
ThreadDict,
|
||||
ThreadFilter,
|
||||
)
|
||||
from chainlit.user import PersistedUser, User
|
||||
|
||||
# Import for runtime usage (isinstance checks)
|
||||
try:
|
||||
from chainlit.data.storage_clients.gcs import GCSStorageClient
|
||||
except ImportError:
|
||||
GCSStorageClient = None # type: ignore[assignment,misc]
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from chainlit.data.storage_clients.gcs import GCSStorageClient
|
||||
from chainlit.element import Element, ElementDict
|
||||
from chainlit.step import StepDict
|
||||
|
||||
ISO_FORMAT = "%Y-%m-%dT%H:%M:%S.%fZ"
|
||||
|
||||
|
||||
class ChainlitDataLayer(BaseDataLayer):
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
storage_client: Optional[BaseStorageClient] = None,
|
||||
show_logger: bool = False,
|
||||
):
|
||||
self.database_url = database_url
|
||||
self.pool: Optional[asyncpg.Pool] = None
|
||||
self.storage_client = storage_client
|
||||
self.show_logger = show_logger
|
||||
|
||||
async def connect(self):
|
||||
if not self.pool:
|
||||
self.pool = await asyncpg.create_pool(self.database_url)
|
||||
|
||||
async def get_current_timestamp(self) -> datetime:
|
||||
return datetime.now()
|
||||
|
||||
async def execute_query(
|
||||
self, query: str, params: Union[Dict, None] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
if not self.pool:
|
||||
await self.connect()
|
||||
|
||||
try:
|
||||
async with self.pool.acquire() as connection: # type: ignore
|
||||
try:
|
||||
if params:
|
||||
records = await connection.fetch(query, *params.values())
|
||||
else:
|
||||
records = await connection.fetch(query)
|
||||
return [dict(record) for record in records]
|
||||
except Exception as e:
|
||||
logger.error(f"Database error: {e!s}")
|
||||
raise
|
||||
except (
|
||||
asyncpg.exceptions.ConnectionDoesNotExistError,
|
||||
asyncpg.exceptions.InterfaceError,
|
||||
) as e:
|
||||
# Handle connection issues by cleaning up and rethrowing
|
||||
logger.error(f"Connection error: {e!s}")
|
||||
await self.cleanup()
|
||||
raise
|
||||
|
||||
async def get_user(self, identifier: str) -> Optional[PersistedUser]:
|
||||
query = """
|
||||
SELECT * FROM "User"
|
||||
WHERE identifier = $1
|
||||
"""
|
||||
result = await self.execute_query(query, {"identifier": identifier})
|
||||
if not result or len(result) == 0:
|
||||
return None
|
||||
row = result[0]
|
||||
|
||||
return PersistedUser(
|
||||
id=str(row.get("id")),
|
||||
identifier=str(row.get("identifier")),
|
||||
createdAt=row.get("createdAt").isoformat(), # type: ignore
|
||||
metadata=json.loads(row.get("metadata", "{}")),
|
||||
)
|
||||
|
||||
async def create_user(self, user: User) -> Optional[PersistedUser]:
|
||||
query = """
|
||||
INSERT INTO "User" (id, identifier, metadata, "createdAt", "updatedAt")
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (identifier) DO UPDATE
|
||||
SET metadata = $3
|
||||
RETURNING *
|
||||
"""
|
||||
now = await self.get_current_timestamp()
|
||||
params = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"identifier": user.identifier,
|
||||
"metadata": json.dumps(user.metadata),
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
result = await self.execute_query(query, params)
|
||||
row = result[0]
|
||||
|
||||
return PersistedUser(
|
||||
id=str(row.get("id")),
|
||||
identifier=str(row.get("identifier")),
|
||||
createdAt=row.get("createdAt").isoformat(), # type: ignore
|
||||
metadata=json.loads(row.get("metadata", "{}")),
|
||||
)
|
||||
|
||||
async def delete_feedback(self, feedback_id: str) -> bool:
|
||||
query = """
|
||||
DELETE FROM "Feedback" WHERE id = $1
|
||||
"""
|
||||
await self.execute_query(query, {"feedback_id": feedback_id})
|
||||
return True
|
||||
|
||||
async def upsert_feedback(self, feedback: Feedback) -> str:
|
||||
query = """
|
||||
INSERT INTO "Feedback" (id, "stepId", name, value, comment)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (id) DO UPDATE
|
||||
SET value = $4, comment = $5
|
||||
RETURNING id
|
||||
"""
|
||||
feedback_id = feedback.id or str(uuid.uuid4())
|
||||
params = {
|
||||
"id": feedback_id,
|
||||
"step_id": feedback.forId,
|
||||
"name": "user_feedback",
|
||||
"value": float(feedback.value),
|
||||
"comment": feedback.comment,
|
||||
}
|
||||
results = await self.execute_query(query, params)
|
||||
return str(results[0]["id"])
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_element(self, element: "Element"):
|
||||
if not element.for_id:
|
||||
return
|
||||
|
||||
if element.thread_id:
|
||||
query = 'SELECT id FROM "Thread" WHERE id = $1'
|
||||
results = await self.execute_query(query, {"thread_id": element.thread_id})
|
||||
if not results:
|
||||
await self.update_thread(thread_id=element.thread_id)
|
||||
|
||||
if element.for_id:
|
||||
query = 'SELECT id FROM "Step" WHERE id = $1'
|
||||
results = await self.execute_query(query, {"step_id": element.for_id})
|
||||
if not results:
|
||||
await self.create_step(
|
||||
{
|
||||
"id": element.for_id,
|
||||
"metadata": {},
|
||||
"type": "run",
|
||||
"start_time": await self.get_current_timestamp(),
|
||||
"end_time": await self.get_current_timestamp(),
|
||||
}
|
||||
)
|
||||
|
||||
# Handle file uploads only if storage_client is configured
|
||||
path = None
|
||||
if self.storage_client:
|
||||
content: Optional[Union[bytes, str]] = None
|
||||
|
||||
if element.path:
|
||||
async with aiofiles.open(element.path, "rb") as f:
|
||||
content = await f.read()
|
||||
elif element.content:
|
||||
content = element.content
|
||||
elif not element.url:
|
||||
raise ValueError("Element url, path or content must be provided")
|
||||
|
||||
if content is not None:
|
||||
if element.thread_id:
|
||||
path = f"threads/{element.thread_id}/files/{element.id}"
|
||||
else:
|
||||
path = f"files/{element.id}"
|
||||
|
||||
content_disposition = (
|
||||
f'attachment; filename="{element.name}"'
|
||||
if not (
|
||||
GCSStorageClient is not None
|
||||
and isinstance(self.storage_client, GCSStorageClient)
|
||||
)
|
||||
else None
|
||||
)
|
||||
await self.storage_client.upload_file(
|
||||
object_key=path,
|
||||
data=content,
|
||||
mime=element.mime or "application/octet-stream",
|
||||
overwrite=True,
|
||||
content_disposition=content_disposition,
|
||||
)
|
||||
|
||||
else:
|
||||
# Log warning only if element has file content that needs uploading
|
||||
if element.path or element.url or element.content:
|
||||
logger.warning(
|
||||
"Data Layer: No storage client configured. "
|
||||
"File will not be uploaded."
|
||||
)
|
||||
|
||||
# Always persist element metadata to database
|
||||
query = """
|
||||
INSERT INTO "Element" (
|
||||
id, "threadId", "stepId", metadata, mime, name, "objectKey", url,
|
||||
"chainlitKey", display, size, language, page, props
|
||||
) VALUES (
|
||||
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14
|
||||
)
|
||||
ON CONFLICT (id) DO UPDATE SET
|
||||
props = EXCLUDED.props
|
||||
"""
|
||||
params = {
|
||||
"id": element.id,
|
||||
"thread_id": element.thread_id,
|
||||
"step_id": element.for_id,
|
||||
"metadata": json.dumps(
|
||||
{
|
||||
"size": element.size,
|
||||
"language": element.language,
|
||||
"display": element.display,
|
||||
"type": element.type,
|
||||
"page": getattr(element, "page", None),
|
||||
}
|
||||
),
|
||||
"mime": element.mime,
|
||||
"name": element.name,
|
||||
"object_key": path,
|
||||
"url": element.url,
|
||||
"chainlit_key": element.chainlit_key,
|
||||
"display": element.display,
|
||||
"size": element.size,
|
||||
"language": element.language,
|
||||
"page": getattr(element, "page", None),
|
||||
"props": json.dumps(getattr(element, "props", {})),
|
||||
}
|
||||
await self.execute_query(query, params)
|
||||
|
||||
async def get_element(
|
||||
self, thread_id: str, element_id: str
|
||||
) -> Optional[ElementDict]:
|
||||
query = """
|
||||
SELECT * FROM "Element"
|
||||
WHERE id = $1 AND "threadId" = $2
|
||||
"""
|
||||
results = await self.execute_query(
|
||||
query, {"element_id": element_id, "thread_id": thread_id}
|
||||
)
|
||||
|
||||
if not results:
|
||||
return None
|
||||
|
||||
row = results[0]
|
||||
metadata = json.loads(row.get("metadata", "{}"))
|
||||
|
||||
return ElementDict(
|
||||
id=str(row["id"]),
|
||||
threadId=str(row["threadId"]),
|
||||
type=metadata.get("type", "file"),
|
||||
url=str(row["url"]),
|
||||
name=str(row["name"]),
|
||||
mime=str(row["mime"]),
|
||||
objectKey=str(row["objectKey"]),
|
||||
forId=str(row["stepId"]),
|
||||
chainlitKey=row.get("chainlitKey"),
|
||||
display=row["display"],
|
||||
size=row["size"],
|
||||
language=row["language"],
|
||||
page=row["page"],
|
||||
autoPlay=row.get("autoPlay"),
|
||||
playerConfig=row.get("playerConfig"),
|
||||
props=json.loads(row.get("props", "{}")),
|
||||
)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_element(self, element_id: str, thread_id: Optional[str] = None):
|
||||
query = """
|
||||
SELECT * FROM "Element"
|
||||
WHERE id = $1
|
||||
"""
|
||||
elements = await self.execute_query(query, {"id": element_id})
|
||||
|
||||
if self.storage_client is not None and len(elements) > 0:
|
||||
if elements[0]["objectKey"]:
|
||||
await self.storage_client.delete_file(
|
||||
object_key=elements[0]["objectKey"]
|
||||
)
|
||||
query = """
|
||||
DELETE FROM "Element"
|
||||
WHERE id = $1
|
||||
"""
|
||||
params = {"id": element_id}
|
||||
|
||||
if thread_id:
|
||||
query += ' AND "threadId" = $2'
|
||||
params["thread_id"] = thread_id
|
||||
|
||||
await self.execute_query(query, params)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_step(self, step_dict: StepDict):
|
||||
if step_dict.get("threadId"):
|
||||
thread_query = 'SELECT id FROM "Thread" WHERE id = $1'
|
||||
thread_results = await self.execute_query(
|
||||
thread_query, {"thread_id": step_dict["threadId"]}
|
||||
)
|
||||
if not thread_results:
|
||||
await self.update_thread(thread_id=step_dict["threadId"])
|
||||
|
||||
if step_dict.get("parentId"):
|
||||
parent_query = 'SELECT id FROM "Step" WHERE id = $1'
|
||||
parent_results = await self.execute_query(
|
||||
parent_query, {"parent_id": step_dict["parentId"]}
|
||||
)
|
||||
if not parent_results:
|
||||
await self.create_step(
|
||||
{
|
||||
"id": step_dict["parentId"],
|
||||
"metadata": {},
|
||||
"type": "run",
|
||||
"createdAt": step_dict.get("createdAt"),
|
||||
}
|
||||
)
|
||||
|
||||
query = """
|
||||
INSERT INTO "Step" (
|
||||
id, "threadId", "parentId", input, metadata, name, output,
|
||||
type, "startTime", "endTime", "showInput", "isError"
|
||||
) VALUES (
|
||||
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12
|
||||
)
|
||||
ON CONFLICT (id) DO UPDATE SET
|
||||
"parentId" = COALESCE(EXCLUDED."parentId", "Step"."parentId"),
|
||||
input = COALESCE(NULLIF(EXCLUDED.input, ''), "Step".input),
|
||||
metadata = CASE
|
||||
WHEN EXCLUDED.metadata <> '{}' THEN EXCLUDED.metadata
|
||||
ELSE "Step".metadata
|
||||
END,
|
||||
name = COALESCE(EXCLUDED.name, "Step".name),
|
||||
output = COALESCE(NULLIF(EXCLUDED.output, ''), "Step".output),
|
||||
type = CASE
|
||||
WHEN EXCLUDED.type = 'run' THEN "Step".type
|
||||
ELSE EXCLUDED.type
|
||||
END,
|
||||
"threadId" = COALESCE(EXCLUDED."threadId", "Step"."threadId"),
|
||||
"endTime" = COALESCE(EXCLUDED."endTime", "Step"."endTime"),
|
||||
"startTime" = LEAST(EXCLUDED."startTime", "Step"."startTime"),
|
||||
"showInput" = COALESCE(EXCLUDED."showInput", "Step"."showInput"),
|
||||
"isError" = COALESCE(EXCLUDED."isError", "Step"."isError")
|
||||
"""
|
||||
|
||||
timestamp = await self.get_current_timestamp()
|
||||
created_at = step_dict.get("createdAt")
|
||||
if created_at:
|
||||
timestamp = datetime.strptime(created_at, ISO_FORMAT)
|
||||
|
||||
params = {
|
||||
"id": step_dict["id"],
|
||||
"thread_id": step_dict.get("threadId"),
|
||||
"parent_id": step_dict.get("parentId"),
|
||||
"input": step_dict.get("input"),
|
||||
"metadata": json.dumps(step_dict.get("metadata", {})),
|
||||
"name": step_dict.get("name"),
|
||||
"output": step_dict.get("output"),
|
||||
"type": step_dict["type"],
|
||||
"start_time": timestamp,
|
||||
"end_time": timestamp,
|
||||
"show_input": str(step_dict.get("showInput", "json")),
|
||||
"is_error": step_dict.get("isError", False),
|
||||
}
|
||||
await self.execute_query(query, params)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def update_step(self, step_dict: StepDict):
|
||||
await self.create_step(step_dict)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_step(self, step_id: str):
|
||||
# Delete associated elements and feedbacks first
|
||||
await self.execute_query(
|
||||
'DELETE FROM "Element" WHERE "stepId" = $1', {"step_id": step_id}
|
||||
)
|
||||
await self.execute_query(
|
||||
'DELETE FROM "Feedback" WHERE "stepId" = $1', {"step_id": step_id}
|
||||
)
|
||||
# Delete the step
|
||||
await self.execute_query(
|
||||
'DELETE FROM "Step" WHERE id = $1', {"step_id": step_id}
|
||||
)
|
||||
|
||||
async def get_step(self, step_id: str) -> Optional[StepDict]:
|
||||
# Get step and related feedback
|
||||
query = """
|
||||
SELECT s.*,
|
||||
f.id feedback_id,
|
||||
f.value feedback_value,
|
||||
f."comment" feedback_comment
|
||||
FROM "Step" s left join "Feedback" f on s.id = f."stepId"
|
||||
WHERE s.id = $1
|
||||
"""
|
||||
result = await self.execute_query(query, {"step_id": step_id})
|
||||
if not result:
|
||||
return None
|
||||
return self._convert_step_row_to_dict(result[0])
|
||||
|
||||
async def get_thread_author(self, thread_id: str) -> str:
|
||||
query = """
|
||||
SELECT u.identifier
|
||||
FROM "Thread" t
|
||||
JOIN "User" u ON t."userId" = u.id
|
||||
WHERE t.id = $1
|
||||
"""
|
||||
results = await self.execute_query(query, {"thread_id": thread_id})
|
||||
if not results:
|
||||
raise ValueError(f"Thread {thread_id} not found")
|
||||
return results[0]["identifier"]
|
||||
|
||||
async def delete_thread(self, thread_id: str):
|
||||
elements_query = """
|
||||
SELECT * FROM "Element"
|
||||
WHERE "threadId" = $1
|
||||
"""
|
||||
elements_results = await self.execute_query(
|
||||
elements_query, {"thread_id": thread_id}
|
||||
)
|
||||
|
||||
if self.storage_client is not None:
|
||||
for elem in elements_results:
|
||||
if elem["objectKey"]:
|
||||
await self.storage_client.delete_file(object_key=elem["objectKey"])
|
||||
|
||||
await self.execute_query(
|
||||
'DELETE FROM "Thread" WHERE id = $1', {"thread_id": thread_id}
|
||||
)
|
||||
|
||||
async def list_threads(
|
||||
self, pagination: Pagination, filters: ThreadFilter
|
||||
) -> PaginatedResponse[ThreadDict]:
|
||||
query = """
|
||||
SELECT
|
||||
t.*,
|
||||
u.identifier as user_identifier,
|
||||
(SELECT COUNT(*) FROM "Thread" WHERE "userId" = t."userId") as total
|
||||
FROM "Thread" t
|
||||
LEFT JOIN "User" u ON t."userId" = u.id
|
||||
WHERE t."deletedAt" IS NULL
|
||||
"""
|
||||
params: Dict[str, Any] = {}
|
||||
param_count = 1
|
||||
|
||||
if filters.search:
|
||||
query += f" AND t.name ILIKE ${param_count}"
|
||||
params["name"] = f"%{filters.search}%"
|
||||
param_count += 1
|
||||
|
||||
if filters.userId:
|
||||
query += f' AND t."userId" = ${param_count}'
|
||||
params["user_id"] = filters.userId
|
||||
param_count += 1
|
||||
|
||||
if pagination.cursor:
|
||||
query += f' AND t."updatedAt" < (SELECT "updatedAt" FROM "Thread" WHERE id = ${param_count})'
|
||||
params["cursor"] = pagination.cursor
|
||||
param_count += 1
|
||||
|
||||
query += f' ORDER BY t."updatedAt" DESC LIMIT ${param_count}'
|
||||
params["limit"] = pagination.first + 1
|
||||
|
||||
results = await self.execute_query(query, params)
|
||||
threads = results
|
||||
|
||||
has_next_page = len(threads) > pagination.first
|
||||
if has_next_page:
|
||||
threads = threads[:-1]
|
||||
|
||||
thread_dicts = []
|
||||
for thread in threads:
|
||||
thread_dict = ThreadDict(
|
||||
id=str(thread["id"]),
|
||||
createdAt=thread["updatedAt"].isoformat(),
|
||||
name=thread["name"],
|
||||
userId=str(thread["userId"]) if thread["userId"] else None,
|
||||
userIdentifier=thread["user_identifier"],
|
||||
metadata=json.loads(thread["metadata"]),
|
||||
steps=[],
|
||||
elements=[],
|
||||
tags=[],
|
||||
)
|
||||
thread_dicts.append(thread_dict)
|
||||
|
||||
return PaginatedResponse(
|
||||
pageInfo=PageInfo(
|
||||
hasNextPage=has_next_page,
|
||||
startCursor=thread_dicts[0]["id"] if thread_dicts else None,
|
||||
endCursor=thread_dicts[-1]["id"] if thread_dicts else None,
|
||||
),
|
||||
data=thread_dicts,
|
||||
)
|
||||
|
||||
async def get_thread(self, thread_id: str) -> Optional[ThreadDict]:
|
||||
query = """
|
||||
SELECT t.*, u.identifier as user_identifier
|
||||
FROM "Thread" t
|
||||
LEFT JOIN "User" u ON t."userId" = u.id
|
||||
WHERE t.id = $1 AND t."deletedAt" IS NULL
|
||||
"""
|
||||
results = await self.execute_query(query, {"thread_id": thread_id})
|
||||
|
||||
if not results:
|
||||
return None
|
||||
|
||||
thread = results[0]
|
||||
|
||||
# Get steps and related feedback
|
||||
steps_query = """
|
||||
SELECT s.*,
|
||||
f.id feedback_id,
|
||||
f.value feedback_value,
|
||||
f."comment" feedback_comment
|
||||
FROM "Step" s left join "Feedback" f on s.id = f."stepId"
|
||||
WHERE s."threadId" = $1
|
||||
ORDER BY "startTime"
|
||||
"""
|
||||
steps_results = await self.execute_query(steps_query, {"thread_id": thread_id})
|
||||
|
||||
# Get elements
|
||||
elements_query = """
|
||||
SELECT * FROM "Element"
|
||||
WHERE "threadId" = $1
|
||||
"""
|
||||
elements_results = await self.execute_query(
|
||||
elements_query, {"thread_id": thread_id}
|
||||
)
|
||||
|
||||
if self.storage_client is not None:
|
||||
for elem in elements_results:
|
||||
if not elem["url"] and elem["objectKey"]:
|
||||
elem["url"] = await self.storage_client.get_read_url(
|
||||
object_key=elem["objectKey"],
|
||||
)
|
||||
|
||||
return ThreadDict(
|
||||
id=str(thread["id"]),
|
||||
createdAt=thread["createdAt"].isoformat(),
|
||||
name=thread["name"],
|
||||
userId=str(thread["userId"]) if thread["userId"] else None,
|
||||
userIdentifier=thread["user_identifier"],
|
||||
metadata=json.loads(thread["metadata"]),
|
||||
steps=[self._convert_step_row_to_dict(step) for step in steps_results],
|
||||
elements=[
|
||||
self._convert_element_row_to_dict(elem) for elem in elements_results
|
||||
],
|
||||
tags=[],
|
||||
)
|
||||
|
||||
async def update_thread(
|
||||
self,
|
||||
thread_id: str,
|
||||
name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
metadata: Optional[Dict] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
if self.show_logger:
|
||||
logger.info(f"asyncpg: update_thread, thread_id={thread_id}")
|
||||
|
||||
has_updates = (
|
||||
metadata is not None
|
||||
or name is not None
|
||||
or user_id is not None
|
||||
or tags is not None
|
||||
)
|
||||
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
|
||||
thread_name = truncate(
|
||||
name
|
||||
if name is not None
|
||||
else (metadata.get("name") if metadata and "name" in metadata else None)
|
||||
)
|
||||
|
||||
existing = await self.execute_query(
|
||||
'SELECT "metadata" FROM "Thread" WHERE id = $1',
|
||||
{"thread_id": thread_id},
|
||||
)
|
||||
|
||||
thread_exists = isinstance(existing, list) and existing
|
||||
if thread_exists and not has_updates:
|
||||
return
|
||||
|
||||
base = {}
|
||||
if thread_exists:
|
||||
raw = existing[0].get("metadata") or {}
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
base = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
base = {}
|
||||
elif isinstance(raw, dict):
|
||||
base = raw
|
||||
to_delete = {k for k, v in metadata.items() if v is None}
|
||||
incoming = {k: v for k, v in metadata.items() if v is not None}
|
||||
base = {k: v for k, v in base.items() if k not in to_delete}
|
||||
metadata = {**base, **incoming}
|
||||
|
||||
data = {
|
||||
"id": thread_id,
|
||||
"name": thread_name,
|
||||
"userId": user_id,
|
||||
"tags": tags,
|
||||
"metadata": json.dumps(metadata),
|
||||
"updatedAt": datetime.now(),
|
||||
}
|
||||
|
||||
# Remove None values
|
||||
data = {k: v for k, v in data.items() if v is not None}
|
||||
|
||||
# Build the query dynamically based on available fields
|
||||
columns = [f'"{k}"' for k in data.keys()]
|
||||
placeholders = [f"${i + 1}" for i in range(len(data))]
|
||||
values = list(data.values())
|
||||
|
||||
update_sets = [f'"{k}" = EXCLUDED."{k}"' for k in data.keys() if k != "id"]
|
||||
|
||||
if update_sets:
|
||||
query = f"""
|
||||
INSERT INTO "Thread" ({", ".join(columns)})
|
||||
VALUES ({", ".join(placeholders)})
|
||||
ON CONFLICT (id) DO UPDATE
|
||||
SET {", ".join(update_sets)};
|
||||
"""
|
||||
else:
|
||||
query = f"""
|
||||
INSERT INTO "Thread" ({", ".join(columns)})
|
||||
VALUES ({", ".join(placeholders)})
|
||||
ON CONFLICT (id) DO NOTHING
|
||||
"""
|
||||
|
||||
await self.execute_query(query, {str(i + 1): v for i, v in enumerate(values)})
|
||||
|
||||
async def get_favorite_steps(self, user_id: str) -> List[StepDict]:
|
||||
query = """
|
||||
SELECT s.*
|
||||
FROM "Step" s
|
||||
JOIN "Thread" t ON s."threadId" = t.id
|
||||
WHERE t."userId" = $1
|
||||
AND s.metadata::jsonb->>'favorite' = 'true'
|
||||
ORDER BY s."createdAt" DESC \
|
||||
"""
|
||||
results = await self.execute_query(query, {"user_id": user_id})
|
||||
return [self._convert_step_row_to_dict(row) for row in results]
|
||||
|
||||
def _extract_feedback_dict_from_step_row(self, row: Dict) -> Optional[FeedbackDict]:
|
||||
if row.get("feedback_id", None) is not None:
|
||||
return FeedbackDict(
|
||||
forId=str(row["id"]),
|
||||
id=str(row["feedback_id"]),
|
||||
value=row["feedback_value"],
|
||||
comment=row["feedback_comment"],
|
||||
)
|
||||
return None
|
||||
|
||||
def _convert_step_row_to_dict(self, row: Dict) -> StepDict:
|
||||
return StepDict(
|
||||
id=str(row["id"]),
|
||||
threadId=str(row["threadId"]) if row.get("threadId") else "",
|
||||
parentId=str(row["parentId"]) if row.get("parentId") else None,
|
||||
name=str(row.get("name")),
|
||||
type=row["type"],
|
||||
input=row.get("input", {}),
|
||||
output=row.get("output", {}),
|
||||
metadata=json.loads(row.get("metadata", "{}")),
|
||||
createdAt=row["createdAt"].isoformat() if row.get("createdAt") else None,
|
||||
start=row["startTime"].isoformat() if row.get("startTime") else None,
|
||||
showInput=row.get("showInput"),
|
||||
isError=row.get("isError"),
|
||||
end=row["endTime"].isoformat() if row.get("endTime") else None,
|
||||
feedback=self._extract_feedback_dict_from_step_row(row),
|
||||
)
|
||||
|
||||
def _convert_element_row_to_dict(self, row: Dict) -> ElementDict:
|
||||
metadata = json.loads(row.get("metadata", "{}"))
|
||||
return ElementDict(
|
||||
id=str(row["id"]),
|
||||
threadId=str(row["threadId"]) if row.get("threadId") else None,
|
||||
type=metadata.get("type", "file"),
|
||||
url=row["url"],
|
||||
name=row["name"],
|
||||
mime=row["mime"],
|
||||
objectKey=row["objectKey"],
|
||||
forId=str(row["stepId"]),
|
||||
chainlitKey=row.get("chainlitKey"),
|
||||
display=row["display"],
|
||||
size=row["size"],
|
||||
language=row["language"],
|
||||
page=row["page"],
|
||||
autoPlay=row.get("autoPlay"),
|
||||
playerConfig=row.get("playerConfig"),
|
||||
props=json.loads(row.get("props") or "{}"),
|
||||
)
|
||||
|
||||
async def build_debug_url(self) -> str:
|
||||
return ""
|
||||
|
||||
async def cleanup(self):
|
||||
"""Cleanup database connections"""
|
||||
if self.pool:
|
||||
logger.debug("Cleaning up connection pool")
|
||||
await self.pool.close()
|
||||
self.pool = None
|
||||
|
||||
async def close(self) -> None:
|
||||
if self.storage_client:
|
||||
await self.storage_client.close()
|
||||
await self.cleanup()
|
||||
|
||||
|
||||
def truncate(text: Optional[str], max_length: int = 255) -> Optional[str]:
|
||||
return None if text is None else text[:max_length]
|
||||
@@ -0,0 +1,687 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
|
||||
|
||||
import aiofiles
|
||||
import aiohttp
|
||||
import boto3 # type: ignore
|
||||
from boto3.dynamodb.types import TypeDeserializer, TypeSerializer
|
||||
|
||||
from chainlit.context import context
|
||||
from chainlit.data.base import BaseDataLayer
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient
|
||||
from chainlit.data.utils import queue_until_user_message
|
||||
from chainlit.element import ElementDict
|
||||
from chainlit.logger import logger
|
||||
from chainlit.step import StepDict
|
||||
from chainlit.types import (
|
||||
Feedback,
|
||||
PageInfo,
|
||||
PaginatedResponse,
|
||||
Pagination,
|
||||
ThreadDict,
|
||||
ThreadFilter,
|
||||
)
|
||||
from chainlit.user import PersistedUser, User
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mypy_boto3_dynamodb import DynamoDBClient
|
||||
|
||||
from chainlit.element import Element
|
||||
|
||||
|
||||
_logger = logger.getChild("DynamoDB")
|
||||
_logger.setLevel(logging.WARNING)
|
||||
|
||||
|
||||
class DynamoDBDataLayer(BaseDataLayer):
|
||||
def __init__(
|
||||
self,
|
||||
table_name: str,
|
||||
client: Optional["DynamoDBClient"] = None,
|
||||
storage_provider: Optional[BaseStorageClient] = None,
|
||||
user_thread_limit: int = 10,
|
||||
):
|
||||
if client:
|
||||
self.client = client
|
||||
else:
|
||||
region_name = os.environ.get("AWS_REGION", "us-east-1")
|
||||
self.client = boto3.client("dynamodb", region_name=region_name) # type: ignore
|
||||
|
||||
self.table_name = table_name
|
||||
self.storage_provider = storage_provider
|
||||
self.user_thread_limit = user_thread_limit
|
||||
|
||||
self._type_deserializer = TypeDeserializer()
|
||||
self._type_serializer = TypeSerializer()
|
||||
|
||||
def _get_current_timestamp(self) -> str:
|
||||
return datetime.now().isoformat() + "Z"
|
||||
|
||||
def _serialize_item(self, item: dict[str, Any]) -> dict[str, Any]:
|
||||
def convert_floats(obj):
|
||||
if isinstance(obj, float):
|
||||
return Decimal(str(obj))
|
||||
elif isinstance(obj, dict):
|
||||
return {k: convert_floats(v) for k, v in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [convert_floats(v) for v in obj]
|
||||
else:
|
||||
return obj
|
||||
|
||||
return {
|
||||
key: self._type_serializer.serialize(convert_floats(value))
|
||||
for key, value in item.items()
|
||||
}
|
||||
|
||||
def _deserialize_item(self, item: dict[str, Any]) -> dict[str, Any]:
|
||||
def convert_decimals(obj):
|
||||
if isinstance(obj, Decimal):
|
||||
return float(obj)
|
||||
elif isinstance(obj, dict):
|
||||
return {k: convert_decimals(v) for k, v in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [convert_decimals(v) for v in obj]
|
||||
else:
|
||||
return obj
|
||||
|
||||
return {
|
||||
key: convert_decimals(self._type_deserializer.deserialize(value))
|
||||
for key, value in item.items()
|
||||
}
|
||||
|
||||
def _update_item(self, key: Dict[str, Any], updates: Dict[str, Any]):
|
||||
update_expr: List[str] = []
|
||||
expression_attribute_names = {}
|
||||
expression_attribute_values = {}
|
||||
|
||||
for index, (attr, value) in enumerate(updates.items()):
|
||||
if value is None:
|
||||
continue
|
||||
|
||||
k, v = f"#{index}", f":{index}"
|
||||
update_expr.append(f"{k} = {v}")
|
||||
expression_attribute_names[k] = attr
|
||||
expression_attribute_values[v] = value
|
||||
|
||||
self.client.update_item(
|
||||
TableName=self.table_name,
|
||||
Key=self._serialize_item(key),
|
||||
UpdateExpression="SET " + ", ".join(update_expr),
|
||||
ExpressionAttributeNames=expression_attribute_names,
|
||||
ExpressionAttributeValues=self._serialize_item(expression_attribute_values),
|
||||
)
|
||||
|
||||
@property
|
||||
def context(self):
|
||||
return context
|
||||
|
||||
async def get_user(self, identifier: str) -> Optional["PersistedUser"]:
|
||||
_logger.info("DynamoDB: get_user identifier=%s", identifier)
|
||||
|
||||
response = self.client.get_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"USER#{identifier}"},
|
||||
"SK": {"S": "USER"},
|
||||
},
|
||||
)
|
||||
|
||||
if "Item" not in response:
|
||||
return None
|
||||
|
||||
user = self._deserialize_item(response["Item"])
|
||||
|
||||
return PersistedUser(
|
||||
id=user["id"],
|
||||
identifier=user["identifier"],
|
||||
createdAt=user["createdAt"],
|
||||
metadata=user["metadata"],
|
||||
)
|
||||
|
||||
async def create_user(self, user: "User") -> Optional["PersistedUser"]:
|
||||
_logger.info("DynamoDB: create_user user.identifier=%s", user.identifier)
|
||||
|
||||
ts = self._get_current_timestamp()
|
||||
metadata: Dict[Any, Any] = user.metadata # type: ignore
|
||||
|
||||
item = {
|
||||
"PK": f"USER#{user.identifier}",
|
||||
"SK": "USER",
|
||||
"id": user.identifier,
|
||||
"identifier": user.identifier,
|
||||
"metadata": metadata,
|
||||
"createdAt": ts,
|
||||
}
|
||||
|
||||
self.client.put_item(
|
||||
TableName=self.table_name,
|
||||
Item=self._serialize_item(item),
|
||||
)
|
||||
|
||||
return PersistedUser(
|
||||
id=user.identifier,
|
||||
identifier=user.identifier,
|
||||
createdAt=ts,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
async def delete_feedback(self, feedback_id: str) -> bool:
|
||||
_logger.info("DynamoDB: delete_feedback feedback_id=%s", feedback_id)
|
||||
|
||||
# feedback id = THREAD#{thread_id}::STEP#{step_id}
|
||||
thread_id, step_id = feedback_id.split("::")
|
||||
thread_id = thread_id.strip("THREAD#")
|
||||
step_id = step_id.strip("STEP#")
|
||||
|
||||
self.client.update_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": f"STEP#{step_id}"},
|
||||
},
|
||||
UpdateExpression="REMOVE #feedback",
|
||||
ExpressionAttributeNames={"#feedback": "feedback"},
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
async def upsert_feedback(self, feedback: Feedback) -> str:
|
||||
_logger.info(
|
||||
"DynamoDB: upsert_feedback thread=%s step=%s value=%s",
|
||||
feedback.threadId,
|
||||
feedback.forId,
|
||||
feedback.value,
|
||||
)
|
||||
|
||||
if not feedback.forId:
|
||||
raise ValueError(
|
||||
"DynamoDB data layer expects value for feedback.threadId got None"
|
||||
)
|
||||
|
||||
feedback.id = f"THREAD#{feedback.threadId}::STEP#{feedback.forId}"
|
||||
serialized_feedback = self._type_serializer.serialize(asdict(feedback))
|
||||
|
||||
self.client.update_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{feedback.threadId}"},
|
||||
"SK": {"S": f"STEP#{feedback.forId}"},
|
||||
},
|
||||
UpdateExpression="SET #feedback = :feedback",
|
||||
ExpressionAttributeNames={"#feedback": "feedback"},
|
||||
ExpressionAttributeValues={":feedback": serialized_feedback},
|
||||
)
|
||||
|
||||
return feedback.id
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_element(self, element: "Element"):
|
||||
_logger.info(
|
||||
"DynamoDB: create_element thread=%s step=%s type=%s",
|
||||
element.thread_id,
|
||||
element.for_id,
|
||||
element.type,
|
||||
)
|
||||
_logger.debug("DynamoDB: create_element: %s", element.to_dict())
|
||||
|
||||
if not element.for_id:
|
||||
return
|
||||
|
||||
if not self.storage_provider:
|
||||
_logger.warning(
|
||||
"DynamoDB: create_element error. No storage_provider is configured!"
|
||||
)
|
||||
return
|
||||
|
||||
content: Optional[Union[bytes, str]] = None
|
||||
|
||||
if element.content:
|
||||
content = element.content
|
||||
|
||||
elif element.path:
|
||||
_logger.debug("DynamoDB: create_element reading file %s", element.path)
|
||||
async with aiofiles.open(element.path, "rb") as f:
|
||||
content = await f.read()
|
||||
|
||||
elif element.url:
|
||||
_logger.debug("DynamoDB: create_element http %s", element.url)
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(element.url) as response:
|
||||
if response.status == 200:
|
||||
content = await response.read()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Failed to read content from {element.url} status {response.status}",
|
||||
)
|
||||
|
||||
else:
|
||||
raise ValueError("Element url, path or content must be provided")
|
||||
|
||||
if content is None:
|
||||
raise ValueError("Content is None, cannot upload file")
|
||||
|
||||
if not element.mime:
|
||||
element.mime = "application/octet-stream"
|
||||
|
||||
context_user = self.context.session.user
|
||||
user_folder = getattr(context_user, "id", "unknown")
|
||||
file_object_key = f"{user_folder}/{element.thread_id}/{element.id}"
|
||||
|
||||
uploaded_file = await self.storage_provider.upload_file(
|
||||
object_key=file_object_key,
|
||||
data=content,
|
||||
mime=element.mime,
|
||||
overwrite=True,
|
||||
)
|
||||
if not uploaded_file:
|
||||
raise ValueError(
|
||||
"DynamoDB Error: create_element, Failed to persist data in storage_provider",
|
||||
)
|
||||
|
||||
element_dict: Dict[str, Any] = element.to_dict() # type: ignore
|
||||
element_dict.update(
|
||||
{
|
||||
"PK": f"THREAD#{element.thread_id}",
|
||||
"SK": f"ELEMENT#{element.id}",
|
||||
"url": uploaded_file.get("url"),
|
||||
"objectKey": uploaded_file.get("object_key"),
|
||||
}
|
||||
)
|
||||
|
||||
self.client.put_item(
|
||||
TableName=self.table_name,
|
||||
Item=self._serialize_item(element_dict),
|
||||
)
|
||||
|
||||
async def get_element(
|
||||
self, thread_id: str, element_id: str
|
||||
) -> Optional["ElementDict"]:
|
||||
_logger.info(
|
||||
"DynamoDB: get_element thread=%s element=%s", thread_id, element_id
|
||||
)
|
||||
|
||||
response = self.client.get_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": f"ELEMENT#{element_id}"},
|
||||
},
|
||||
)
|
||||
|
||||
if "Item" not in response:
|
||||
return None
|
||||
|
||||
return self._deserialize_item(response["Item"]) # type: ignore
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_element(self, element_id: str, thread_id: Optional[str] = None):
|
||||
thread_id = self.context.session.thread_id
|
||||
_logger.info(
|
||||
"DynamoDB: delete_element thread=%s element=%s", thread_id, element_id
|
||||
)
|
||||
|
||||
self.client.delete_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": f"ELEMENT#{element_id}"},
|
||||
},
|
||||
)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_step(self, step_dict: "StepDict"):
|
||||
_logger.info(
|
||||
"DynamoDB: create_step thread=%s step=%s",
|
||||
step_dict.get("threadId"),
|
||||
step_dict.get("id"),
|
||||
)
|
||||
_logger.debug("DynamoDB: create_step: %s", step_dict)
|
||||
|
||||
item = dict(step_dict)
|
||||
item.update(
|
||||
{
|
||||
# ignore type, dynamo needs these so we want to fail if not set
|
||||
"PK": f"THREAD#{step_dict['threadId']}", # type: ignore
|
||||
"SK": f"STEP#{step_dict['id']}", # type: ignore
|
||||
}
|
||||
)
|
||||
|
||||
self.client.put_item(
|
||||
TableName=self.table_name,
|
||||
Item=self._serialize_item(item),
|
||||
)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def update_step(self, step_dict: "StepDict"):
|
||||
_logger.info(
|
||||
"DynamoDB: update_step thread=%s step=%s",
|
||||
step_dict.get("threadId"),
|
||||
step_dict.get("id"),
|
||||
)
|
||||
_logger.debug("DynamoDB: update_step: %s", step_dict)
|
||||
|
||||
self._update_item(
|
||||
key={
|
||||
# ignore type, dynamo needs these so we want to fail if not set
|
||||
"PK": f"THREAD#{step_dict['threadId']}", # type: ignore
|
||||
"SK": f"STEP#{step_dict['id']}", # type: ignore
|
||||
},
|
||||
updates=step_dict, # type: ignore
|
||||
)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_step(self, step_id: str):
|
||||
thread_id = self.context.session.thread_id
|
||||
_logger.info("DynamoDB: delete_feedback thread=%s step=%s", thread_id, step_id)
|
||||
|
||||
self.client.delete_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": f"STEP#{step_id}"},
|
||||
},
|
||||
)
|
||||
|
||||
async def get_thread_author(self, thread_id: str) -> str:
|
||||
_logger.info("DynamoDB: get_thread_author thread=%s", thread_id)
|
||||
|
||||
response = self.client.get_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": "THREAD"},
|
||||
},
|
||||
ProjectionExpression="userId",
|
||||
)
|
||||
|
||||
if "Item" not in response:
|
||||
raise ValueError(f"Author not found for thread_id {thread_id}")
|
||||
|
||||
item = self._deserialize_item(response["Item"])
|
||||
return item["userId"]
|
||||
|
||||
async def delete_thread(self, thread_id: str):
|
||||
_logger.info("DynamoDB: delete_thread thread=%s", thread_id)
|
||||
|
||||
thread = await self.get_thread(thread_id)
|
||||
if not thread:
|
||||
return
|
||||
|
||||
items: List[Any] = thread["steps"]
|
||||
if thread["elements"]:
|
||||
items.extend(thread["elements"])
|
||||
|
||||
delete_requests = []
|
||||
for item in items:
|
||||
key = self._serialize_item({"PK": item["PK"], "SK": item["SK"]})
|
||||
req = {"DeleteRequest": {"Key": key}}
|
||||
delete_requests.append(req)
|
||||
|
||||
BATCH_ITEM_SIZE = 25 # pylint: disable=invalid-name
|
||||
for i in range(0, len(delete_requests), BATCH_ITEM_SIZE):
|
||||
chunk = delete_requests[i : i + BATCH_ITEM_SIZE]
|
||||
response = self.client.batch_write_item(
|
||||
RequestItems={
|
||||
self.table_name: chunk, # type: ignore
|
||||
}
|
||||
)
|
||||
|
||||
backoff_time = 1
|
||||
while response.get("UnprocessedItems"):
|
||||
backoff_time *= 2
|
||||
# Cap the backoff time at 32 seconds & add jitter
|
||||
delay = min(backoff_time, 32) + random.uniform(0, 1)
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
response = self.client.batch_write_item(
|
||||
RequestItems=response["UnprocessedItems"]
|
||||
)
|
||||
|
||||
self.client.delete_item(
|
||||
TableName=self.table_name,
|
||||
Key={
|
||||
"PK": {"S": f"THREAD#{thread_id}"},
|
||||
"SK": {"S": "THREAD"},
|
||||
},
|
||||
)
|
||||
|
||||
async def list_threads(
|
||||
self, pagination: "Pagination", filters: "ThreadFilter"
|
||||
) -> "PaginatedResponse[ThreadDict]":
|
||||
_logger.info("DynamoDB: list_threads filters.userId=%s", filters.userId)
|
||||
|
||||
if filters.feedback:
|
||||
_logger.warning("DynamoDB: filters on feedback not supported")
|
||||
|
||||
paginated_response: PaginatedResponse[ThreadDict] = PaginatedResponse(
|
||||
data=[],
|
||||
pageInfo=PageInfo(
|
||||
hasNextPage=False, startCursor=pagination.cursor, endCursor=None
|
||||
),
|
||||
)
|
||||
|
||||
query_args: Dict[str, Any] = {
|
||||
"TableName": self.table_name,
|
||||
"IndexName": "UserThread",
|
||||
"ScanIndexForward": False,
|
||||
"Limit": self.user_thread_limit,
|
||||
"KeyConditionExpression": "#UserThreadPK = :pk",
|
||||
"ExpressionAttributeNames": {
|
||||
"#UserThreadPK": "UserThreadPK",
|
||||
},
|
||||
"ExpressionAttributeValues": {
|
||||
":pk": {"S": f"USER#{filters.userId}"},
|
||||
},
|
||||
}
|
||||
|
||||
if pagination.cursor:
|
||||
query_args["ExclusiveStartKey"] = json.loads(pagination.cursor)
|
||||
|
||||
if filters.search:
|
||||
query_args["FilterExpression"] = "contains(#name, :search)"
|
||||
query_args["ExpressionAttributeNames"]["#name"] = "name"
|
||||
query_args["ExpressionAttributeValues"][":search"] = {"S": filters.search}
|
||||
|
||||
response = self.client.query(**query_args) # type: ignore
|
||||
|
||||
if "LastEvaluatedKey" in response:
|
||||
paginated_response.pageInfo.hasNextPage = True
|
||||
paginated_response.pageInfo.endCursor = json.dumps(
|
||||
response["LastEvaluatedKey"]
|
||||
)
|
||||
|
||||
for item in response["Items"]:
|
||||
deserialized_item: Dict[str, Any] = self._deserialize_item(item)
|
||||
thread = ThreadDict( # type: ignore
|
||||
id=deserialized_item["PK"].strip("THREAD#"),
|
||||
createdAt=deserialized_item["UserThreadSK"].strip("TS#"),
|
||||
name=deserialized_item["name"],
|
||||
)
|
||||
paginated_response.data.append(thread)
|
||||
|
||||
return paginated_response
|
||||
|
||||
async def get_thread(self, thread_id: str) -> "Optional[ThreadDict]":
|
||||
_logger.info("DynamoDB: get_thread thread=%s", thread_id)
|
||||
|
||||
# Get all thread records
|
||||
thread_items: List[Any] = []
|
||||
|
||||
cursor: Dict[str, Any] = {}
|
||||
while True:
|
||||
response = self.client.query(
|
||||
TableName=self.table_name,
|
||||
KeyConditionExpression="#pk = :pk",
|
||||
ExpressionAttributeNames={"#pk": "PK"},
|
||||
ExpressionAttributeValues={":pk": {"S": f"THREAD#{thread_id}"}},
|
||||
**cursor,
|
||||
)
|
||||
|
||||
deserialized_items = map(self._deserialize_item, response["Items"])
|
||||
thread_items.extend(deserialized_items)
|
||||
|
||||
if "LastEvaluatedKey" not in response:
|
||||
break
|
||||
cursor["ExclusiveStartKey"] = response["LastEvaluatedKey"]
|
||||
|
||||
if len(thread_items) == 0:
|
||||
return None
|
||||
|
||||
# process accordingly
|
||||
thread_dict: Optional[ThreadDict] = None
|
||||
steps = []
|
||||
elements = []
|
||||
|
||||
for item in thread_items:
|
||||
if item["SK"] == "THREAD":
|
||||
thread_dict = item
|
||||
|
||||
elif item["SK"].startswith("ELEMENT"):
|
||||
if self.storage_provider is not None:
|
||||
item["url"] = await self.storage_provider.get_read_url(
|
||||
object_key=item["objectKey"],
|
||||
)
|
||||
elements.append(item)
|
||||
|
||||
elif item["SK"].startswith("STEP"):
|
||||
if "feedback" in item: # Decimal is not json serializable
|
||||
item["feedback"]["value"] = int(item["feedback"]["value"])
|
||||
steps.append(item)
|
||||
|
||||
if not thread_dict:
|
||||
if len(thread_items) > 0:
|
||||
_logger.warning(
|
||||
"DynamoDB: found orphaned items for thread=%s", thread_id
|
||||
)
|
||||
return None
|
||||
|
||||
steps.sort(key=lambda i: i["createdAt"])
|
||||
thread_dict.update(
|
||||
{
|
||||
"steps": steps,
|
||||
"elements": elements,
|
||||
}
|
||||
)
|
||||
|
||||
return thread_dict
|
||||
|
||||
async def update_thread(
|
||||
self,
|
||||
thread_id: str,
|
||||
name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
metadata: Optional[Dict] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
_logger.info("DynamoDB: update_thread thread=%s userId=%s", thread_id, user_id)
|
||||
_logger.debug(
|
||||
"DynamoDB: update_thread name=%s tags=%s metadata=%s", name, tags, metadata
|
||||
)
|
||||
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
|
||||
ts = self._get_current_timestamp()
|
||||
|
||||
item = {
|
||||
# GSI: UserThread
|
||||
"UserThreadSK": f"TS#{ts}",
|
||||
#
|
||||
"id": thread_id,
|
||||
"createdAt": ts,
|
||||
"name": name,
|
||||
"userId": user_id,
|
||||
"userIdentifier": user_id,
|
||||
"tags": tags,
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
if user_id:
|
||||
# user_id may be None on subsequent calls, don't update UserThreadPK to "USER#{None}"
|
||||
item["UserThreadPK"] = f"USER#{user_id}"
|
||||
|
||||
self._update_item(
|
||||
key={
|
||||
"PK": f"THREAD#{thread_id}",
|
||||
"SK": "THREAD",
|
||||
},
|
||||
updates=item,
|
||||
)
|
||||
|
||||
async def get_favorite_steps(self, user_id: str) -> List["StepDict"]:
|
||||
_logger.info("DynamoDB: get_favorite_steps user_id=%s", user_id)
|
||||
|
||||
thread_ids = []
|
||||
query_args: Dict[str, Any] = {
|
||||
"TableName": self.table_name,
|
||||
"IndexName": "UserThread",
|
||||
"KeyConditionExpression": "#UserThreadPK = :pk",
|
||||
"ExpressionAttributeNames": {"#UserThreadPK": "UserThreadPK"},
|
||||
"ExpressionAttributeValues": {":pk": {"S": f"USER#{user_id}"}},
|
||||
}
|
||||
|
||||
while True:
|
||||
response = self.client.query(**query_args) # type: ignore
|
||||
for item in response.get("Items", []):
|
||||
pk = item.get("PK", {}).get("S")
|
||||
if pk:
|
||||
thread_ids.append(pk.removeprefix("THREAD#"))
|
||||
|
||||
if "LastEvaluatedKey" not in response:
|
||||
break
|
||||
query_args["ExclusiveStartKey"] = response["LastEvaluatedKey"]
|
||||
|
||||
favorite_steps: List[Dict[str, Any]] = []
|
||||
|
||||
for thread_id in thread_ids:
|
||||
t_query_args: Dict[str, Any] = {
|
||||
"TableName": self.table_name,
|
||||
"KeyConditionExpression": "#pk = :pk AND begins_with(#sk, :sk_prefix)",
|
||||
"FilterExpression": "#metadata.#favorite = :true",
|
||||
"ExpressionAttributeNames": {
|
||||
"#pk": "PK",
|
||||
"#sk": "SK",
|
||||
"#metadata": "metadata",
|
||||
"#favorite": "favorite",
|
||||
},
|
||||
"ExpressionAttributeValues": {
|
||||
":pk": {"S": f"THREAD#{thread_id}"},
|
||||
":sk_prefix": {"S": "STEP#"},
|
||||
":true": {"BOOL": True},
|
||||
},
|
||||
}
|
||||
|
||||
while True:
|
||||
response = self.client.query(**t_query_args) # type: ignore
|
||||
for item in response.get("Items", []):
|
||||
step = self._deserialize_item(item)
|
||||
if "PK" in step:
|
||||
del step["PK"]
|
||||
if "SK" in step:
|
||||
del step["SK"]
|
||||
if "feedback" in step:
|
||||
del step["feedback"]
|
||||
|
||||
favorite_steps.append(step)
|
||||
|
||||
if "LastEvaluatedKey" not in response:
|
||||
break
|
||||
t_query_args["ExclusiveStartKey"] = response["LastEvaluatedKey"]
|
||||
|
||||
favorite_steps.sort(key=lambda x: x.get("createdAt", ""), reverse=True)
|
||||
return cast(List["StepDict"], favorite_steps)
|
||||
|
||||
async def build_debug_url(self) -> str:
|
||||
return ""
|
||||
|
||||
async def close(self) -> None:
|
||||
if self.storage_provider:
|
||||
await self.storage_provider.close()
|
||||
self.client.close()
|
||||
@@ -0,0 +1,524 @@
|
||||
import json
|
||||
|
||||
# Deprecation warning for users of this provider
|
||||
import sys
|
||||
import warnings
|
||||
from typing import Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
import aiofiles
|
||||
from httpx import HTTPStatusError, RequestError
|
||||
from literalai import (
|
||||
Attachment as LiteralAttachment,
|
||||
Score as LiteralScore,
|
||||
Step as LiteralStep,
|
||||
Thread as LiteralThread,
|
||||
)
|
||||
from literalai.observability.filter import threads_filters as LiteralThreadsFilters
|
||||
from literalai.observability.step import StepDict as LiteralStepDict
|
||||
|
||||
from chainlit.data.base import BaseDataLayer
|
||||
from chainlit.data.utils import queue_until_user_message
|
||||
from chainlit.element import Audio, Element, ElementDict, File, Image, Pdf, Text, Video
|
||||
from chainlit.logger import logger
|
||||
from chainlit.step import (
|
||||
FeedbackDict,
|
||||
Step,
|
||||
StepDict,
|
||||
StepType,
|
||||
TrueStepType,
|
||||
check_add_step_in_cot,
|
||||
stub_step,
|
||||
)
|
||||
from chainlit.types import (
|
||||
Feedback,
|
||||
PageInfo,
|
||||
PaginatedResponse,
|
||||
Pagination,
|
||||
ThreadDict,
|
||||
ThreadFilter,
|
||||
)
|
||||
from chainlit.user import PersistedUser, User
|
||||
|
||||
|
||||
def _show_deprecation_warning():
|
||||
message = (
|
||||
"\n\033[93mWARNING: The LiteralAI data provider is being deprecated and will be turned off on October 31st, 2025.\033[0m\n"
|
||||
"Please migrate your data layer to another provider as soon as possible.\n"
|
||||
)
|
||||
print(message, file=sys.stderr)
|
||||
warnings.warn(message, DeprecationWarning, stacklevel=2)
|
||||
|
||||
|
||||
_show_deprecation_warning()
|
||||
|
||||
|
||||
class LiteralToChainlitConverter:
|
||||
@classmethod
|
||||
def steptype_to_steptype(cls, step_type: Optional[StepType]) -> TrueStepType:
|
||||
return cast(TrueStepType, step_type or "undefined")
|
||||
|
||||
@classmethod
|
||||
def score_to_feedbackdict(
|
||||
cls,
|
||||
score: Optional[LiteralScore],
|
||||
) -> "Optional[FeedbackDict]":
|
||||
if not score:
|
||||
return None
|
||||
return {
|
||||
"id": score.id or "",
|
||||
"forId": score.step_id or "",
|
||||
"value": cast(Literal[0, 1], score.value),
|
||||
"comment": score.comment,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def step_to_stepdict(cls, step: LiteralStep) -> "StepDict":
|
||||
metadata = step.metadata or {}
|
||||
input = (step.input or {}).get("content") or (
|
||||
json.dumps(step.input) if step.input and step.input != {} else ""
|
||||
)
|
||||
output = (step.output or {}).get("content") or (
|
||||
json.dumps(step.output) if step.output and step.output != {} else ""
|
||||
)
|
||||
|
||||
user_feedback = (
|
||||
next(
|
||||
(
|
||||
s
|
||||
for s in step.scores
|
||||
if s.type == "HUMAN" and s.name == "user-feedback"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if step.scores
|
||||
else None
|
||||
)
|
||||
|
||||
return {
|
||||
"createdAt": step.created_at,
|
||||
"id": step.id or "",
|
||||
"threadId": step.thread_id or "",
|
||||
"parentId": step.parent_id,
|
||||
"feedback": cls.score_to_feedbackdict(user_feedback),
|
||||
"start": step.start_time,
|
||||
"end": step.end_time,
|
||||
"type": step.type or "undefined",
|
||||
"name": step.name or "",
|
||||
"generation": step.generation.to_dict() if step.generation else None,
|
||||
"input": input,
|
||||
"output": output,
|
||||
"showInput": metadata.get("showInput", False),
|
||||
"language": metadata.get("language"),
|
||||
"isError": bool(step.error),
|
||||
"waitForAnswer": metadata.get("waitForAnswer", False),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def attachment_to_elementdict(cls, attachment: LiteralAttachment) -> ElementDict:
|
||||
metadata = attachment.metadata or {}
|
||||
return {
|
||||
"chainlitKey": None,
|
||||
"display": metadata.get("display", "side"),
|
||||
"language": metadata.get("language"),
|
||||
"autoPlay": metadata.get("autoPlay", None),
|
||||
"playerConfig": metadata.get("playerConfig", None),
|
||||
"page": metadata.get("page"),
|
||||
"props": metadata.get("props"),
|
||||
"size": metadata.get("size"),
|
||||
"type": metadata.get("type", "file"),
|
||||
"forId": attachment.step_id,
|
||||
"id": attachment.id or "",
|
||||
"mime": attachment.mime,
|
||||
"name": attachment.name or "",
|
||||
"objectKey": attachment.object_key,
|
||||
"url": attachment.url,
|
||||
"threadId": attachment.thread_id,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def attachment_to_element(
|
||||
cls, attachment: LiteralAttachment, thread_id: Optional[str] = None
|
||||
) -> Element:
|
||||
metadata = attachment.metadata or {}
|
||||
element_type = metadata.get("type", "file")
|
||||
|
||||
element_class = {
|
||||
"file": File,
|
||||
"image": Image,
|
||||
"audio": Audio,
|
||||
"video": Video,
|
||||
"text": Text,
|
||||
"pdf": Pdf,
|
||||
}.get(element_type, Element)
|
||||
|
||||
assert thread_id or attachment.thread_id
|
||||
|
||||
element = element_class(
|
||||
name=attachment.name or "",
|
||||
display=metadata.get("display", "side"),
|
||||
language=metadata.get("language"),
|
||||
size=metadata.get("size"),
|
||||
url=attachment.url,
|
||||
mime=attachment.mime,
|
||||
thread_id=thread_id or attachment.thread_id,
|
||||
)
|
||||
element.id = attachment.id or ""
|
||||
element.for_id = attachment.step_id
|
||||
element.object_key = attachment.object_key
|
||||
return element
|
||||
|
||||
@classmethod
|
||||
def step_to_step(cls, step: LiteralStep) -> Step:
|
||||
chainlit_step = Step(
|
||||
name=step.name or "",
|
||||
type=cls.steptype_to_steptype(step.type),
|
||||
id=step.id,
|
||||
parent_id=step.parent_id,
|
||||
thread_id=step.thread_id or None,
|
||||
)
|
||||
chainlit_step.start = step.start_time
|
||||
chainlit_step.end = step.end_time
|
||||
chainlit_step.created_at = step.created_at
|
||||
chainlit_step.input = step.input.get("content", "") if step.input else ""
|
||||
chainlit_step.output = step.output.get("content", "") if step.output else ""
|
||||
chainlit_step.is_error = bool(step.error)
|
||||
chainlit_step.metadata = step.metadata or {}
|
||||
chainlit_step.tags = step.tags
|
||||
chainlit_step.generation = step.generation
|
||||
|
||||
if step.attachments:
|
||||
chainlit_step.elements = [
|
||||
cls.attachment_to_element(attachment, chainlit_step.thread_id)
|
||||
for attachment in step.attachments
|
||||
]
|
||||
|
||||
return chainlit_step
|
||||
|
||||
@classmethod
|
||||
def thread_to_threaddict(cls, thread: LiteralThread) -> ThreadDict:
|
||||
return {
|
||||
"id": thread.id,
|
||||
"createdAt": getattr(thread, "created_at", ""),
|
||||
"name": thread.name,
|
||||
"userId": thread.participant_id,
|
||||
"userIdentifier": thread.participant_identifier,
|
||||
"tags": thread.tags,
|
||||
"metadata": thread.metadata,
|
||||
"steps": [cls.step_to_stepdict(step) for step in thread.steps]
|
||||
if thread.steps
|
||||
else [],
|
||||
"elements": [
|
||||
cls.attachment_to_elementdict(attachment)
|
||||
for step in thread.steps
|
||||
for attachment in step.attachments
|
||||
]
|
||||
if thread.steps
|
||||
else [],
|
||||
}
|
||||
|
||||
|
||||
class LiteralDataLayer(BaseDataLayer):
|
||||
def __init__(self, api_key: str, server: Optional[str]):
|
||||
from literalai import AsyncLiteralClient
|
||||
|
||||
self.client = AsyncLiteralClient(api_key=api_key, url=server)
|
||||
logger.info("Chainlit data layer initialized")
|
||||
|
||||
async def build_debug_url(self) -> str:
|
||||
try:
|
||||
project_id = await self.client.api.get_my_project_id()
|
||||
return f"{self.client.api.url}/projects/{project_id}/logs/threads/[thread_id]?currentStepId=[step_id]"
|
||||
except Exception as e:
|
||||
logger.error(f"Error building debug url: {e}")
|
||||
return ""
|
||||
|
||||
async def get_user(self, identifier: str) -> Optional[PersistedUser]:
|
||||
user = await self.client.api.get_user(identifier=identifier)
|
||||
if not user:
|
||||
return None
|
||||
return PersistedUser(
|
||||
id=user.id or "",
|
||||
identifier=user.identifier or "",
|
||||
metadata=user.metadata,
|
||||
createdAt=user.created_at or "",
|
||||
)
|
||||
|
||||
async def create_user(self, user: User) -> Optional[PersistedUser]:
|
||||
_user = await self.client.api.get_user(identifier=user.identifier)
|
||||
if not _user:
|
||||
_user = await self.client.api.create_user(
|
||||
identifier=user.identifier, metadata=user.metadata
|
||||
)
|
||||
elif _user.id:
|
||||
await self.client.api.update_user(id=_user.id, metadata=user.metadata)
|
||||
return PersistedUser(
|
||||
id=_user.id or "",
|
||||
identifier=_user.identifier or "",
|
||||
metadata=user.metadata,
|
||||
createdAt=_user.created_at or "",
|
||||
)
|
||||
|
||||
async def delete_feedback(
|
||||
self,
|
||||
feedback_id: str,
|
||||
):
|
||||
if feedback_id:
|
||||
await self.client.api.delete_score(
|
||||
id=feedback_id,
|
||||
)
|
||||
return True
|
||||
return False
|
||||
|
||||
async def upsert_feedback(
|
||||
self,
|
||||
feedback: Feedback,
|
||||
):
|
||||
if feedback.id:
|
||||
await self.client.api.update_score(
|
||||
id=feedback.id,
|
||||
update_params={
|
||||
"comment": feedback.comment,
|
||||
"value": feedback.value,
|
||||
},
|
||||
)
|
||||
return feedback.id
|
||||
else:
|
||||
created = await self.client.api.create_score(
|
||||
step_id=feedback.forId,
|
||||
value=feedback.value,
|
||||
comment=feedback.comment,
|
||||
name="user-feedback",
|
||||
type="HUMAN",
|
||||
)
|
||||
return created.id or ""
|
||||
|
||||
async def safely_send_steps(self, steps):
|
||||
try:
|
||||
await self.client.api.send_steps(steps)
|
||||
except HTTPStatusError as e:
|
||||
logger.error(f"HTTP Request: error sending steps: {e.response.status_code}")
|
||||
except RequestError as e:
|
||||
logger.error(f"HTTP Request: error for {e.request.url!r}.")
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_element(self, element: "Element"):
|
||||
metadata = {
|
||||
"size": element.size,
|
||||
"language": element.language,
|
||||
"display": element.display,
|
||||
"type": element.type,
|
||||
"page": getattr(element, "page", None),
|
||||
"props": getattr(element, "props", None),
|
||||
}
|
||||
|
||||
if not element.for_id:
|
||||
return
|
||||
|
||||
object_key = None
|
||||
|
||||
if not element.url:
|
||||
if element.path:
|
||||
async with aiofiles.open(element.path, "rb") as f:
|
||||
content: Union[bytes, str] = await f.read()
|
||||
elif element.content:
|
||||
content = element.content
|
||||
else:
|
||||
raise ValueError("Either path or content must be provided")
|
||||
uploaded = await self.client.api.upload_file(
|
||||
content=content, mime=element.mime, thread_id=element.thread_id
|
||||
)
|
||||
object_key = uploaded["object_key"]
|
||||
|
||||
await self.safely_send_steps(
|
||||
[
|
||||
{
|
||||
"id": element.for_id,
|
||||
"threadId": element.thread_id,
|
||||
"attachments": [
|
||||
{
|
||||
"id": element.id,
|
||||
"name": element.name,
|
||||
"metadata": metadata,
|
||||
"mime": element.mime,
|
||||
"url": element.url,
|
||||
"objectKey": object_key,
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
async def get_element(
|
||||
self, thread_id: str, element_id: str
|
||||
) -> Optional["ElementDict"]:
|
||||
attachment = await self.client.api.get_attachment(id=element_id)
|
||||
if not attachment:
|
||||
return None
|
||||
return LiteralToChainlitConverter.attachment_to_elementdict(attachment)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_element(self, element_id: str, thread_id: Optional[str] = None):
|
||||
await self.client.api.delete_attachment(id=element_id)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_step(self, step_dict: "StepDict"):
|
||||
metadata = dict(
|
||||
step_dict.get("metadata", {}),
|
||||
waitForAnswer=step_dict.get("waitForAnswer"),
|
||||
language=step_dict.get("language"),
|
||||
showInput=step_dict.get("showInput"),
|
||||
)
|
||||
|
||||
step: LiteralStepDict = {
|
||||
"createdAt": step_dict.get("createdAt"),
|
||||
"startTime": step_dict.get("start"),
|
||||
"endTime": step_dict.get("end"),
|
||||
"generation": step_dict.get("generation"),
|
||||
"id": step_dict.get("id"),
|
||||
"parentId": step_dict.get("parentId"),
|
||||
"name": step_dict.get("name"),
|
||||
"threadId": step_dict.get("threadId"),
|
||||
"type": step_dict.get("type"),
|
||||
"tags": step_dict.get("tags"),
|
||||
"metadata": metadata,
|
||||
}
|
||||
if step_dict.get("input"):
|
||||
step["input"] = {"content": step_dict.get("input")}
|
||||
if step_dict.get("output"):
|
||||
step["output"] = {"content": step_dict.get("output")}
|
||||
if step_dict.get("isError"):
|
||||
step["error"] = step_dict.get("output")
|
||||
|
||||
await self.safely_send_steps([step])
|
||||
|
||||
@queue_until_user_message()
|
||||
async def update_step(self, step_dict: "StepDict"):
|
||||
await self.create_step(step_dict)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_step(self, step_id: str):
|
||||
await self.client.api.delete_step(id=step_id)
|
||||
|
||||
async def get_thread_author(self, thread_id: str) -> str:
|
||||
thread = await self.get_thread(thread_id)
|
||||
if not thread:
|
||||
return ""
|
||||
user_identifier = thread.get("userIdentifier")
|
||||
if not user_identifier:
|
||||
return ""
|
||||
|
||||
return user_identifier
|
||||
|
||||
async def delete_thread(self, thread_id: str):
|
||||
await self.client.api.delete_thread(id=thread_id)
|
||||
|
||||
async def list_threads(
|
||||
self, pagination: "Pagination", filters: "ThreadFilter"
|
||||
) -> "PaginatedResponse[ThreadDict]":
|
||||
if not filters.userId:
|
||||
raise ValueError("userId is required")
|
||||
|
||||
literal_filters: LiteralThreadsFilters = [
|
||||
{
|
||||
"field": "participantId",
|
||||
"operator": "eq",
|
||||
"value": filters.userId,
|
||||
}
|
||||
]
|
||||
|
||||
if filters.search:
|
||||
literal_filters.append(
|
||||
{
|
||||
"field": "stepOutput",
|
||||
"operator": "ilike",
|
||||
"value": filters.search,
|
||||
"path": "content",
|
||||
}
|
||||
)
|
||||
|
||||
if filters.feedback is not None:
|
||||
literal_filters.append(
|
||||
{
|
||||
"field": "scoreValue",
|
||||
"operator": "eq",
|
||||
"value": filters.feedback,
|
||||
"path": "user-feedback",
|
||||
}
|
||||
)
|
||||
|
||||
literal_response = await self.client.api.list_threads(
|
||||
first=pagination.first,
|
||||
after=pagination.cursor,
|
||||
filters=literal_filters,
|
||||
order_by={"column": "createdAt", "direction": "DESC"},
|
||||
)
|
||||
|
||||
chainlit_threads = [
|
||||
*map(LiteralToChainlitConverter.thread_to_threaddict, literal_response.data)
|
||||
]
|
||||
|
||||
return PaginatedResponse(
|
||||
pageInfo=PageInfo(
|
||||
hasNextPage=literal_response.page_info.has_next_page,
|
||||
startCursor=literal_response.page_info.start_cursor,
|
||||
endCursor=literal_response.page_info.end_cursor,
|
||||
),
|
||||
data=chainlit_threads,
|
||||
)
|
||||
|
||||
async def get_thread(self, thread_id: str) -> Optional[ThreadDict]:
|
||||
thread = await self.client.api.get_thread(id=thread_id)
|
||||
if not thread:
|
||||
return None
|
||||
|
||||
elements: List[ElementDict] = []
|
||||
steps: List[StepDict] = []
|
||||
if thread.steps:
|
||||
for step in thread.steps:
|
||||
for attachment in step.attachments:
|
||||
elements.append(
|
||||
LiteralToChainlitConverter.attachment_to_elementdict(attachment)
|
||||
)
|
||||
|
||||
chainlit_step = LiteralToChainlitConverter.step_to_step(step)
|
||||
if check_add_step_in_cot(chainlit_step):
|
||||
steps.append(
|
||||
LiteralToChainlitConverter.step_to_stepdict(step)
|
||||
) # TODO: chainlit_step.to_dict()
|
||||
else:
|
||||
steps.append(stub_step(chainlit_step))
|
||||
|
||||
return {
|
||||
"createdAt": thread.created_at or "",
|
||||
"id": thread.id,
|
||||
"name": thread.name or None,
|
||||
"steps": steps,
|
||||
"elements": elements,
|
||||
"metadata": thread.metadata,
|
||||
"userId": thread.participant_id,
|
||||
"userIdentifier": thread.participant_identifier,
|
||||
"tags": thread.tags,
|
||||
}
|
||||
|
||||
async def update_thread(
|
||||
self,
|
||||
thread_id: str,
|
||||
name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
metadata: Optional[Dict] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
await self.client.api.upsert_thread(
|
||||
id=thread_id,
|
||||
name=name,
|
||||
participant_id=user_id,
|
||||
metadata=metadata,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
async def get_favorite_steps(self, user_id: str) -> List[StepDict]:
|
||||
"""noop for literalai"""
|
||||
return []
|
||||
|
||||
async def close(self):
|
||||
self.client.flush_and_stop()
|
||||
@@ -0,0 +1,953 @@
|
||||
import json
|
||||
import ssl
|
||||
import uuid
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
import aiofiles
|
||||
import aiohttp
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, create_async_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from chainlit.data.base import BaseDataLayer
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient
|
||||
from chainlit.data.utils import queue_until_user_message
|
||||
from chainlit.element import ElementDict
|
||||
from chainlit.logger import logger
|
||||
from chainlit.step import StepDict
|
||||
from chainlit.types import (
|
||||
Feedback,
|
||||
FeedbackDict,
|
||||
PageInfo,
|
||||
PaginatedResponse,
|
||||
Pagination,
|
||||
ThreadDict,
|
||||
ThreadFilter,
|
||||
)
|
||||
from chainlit.user import PersistedUser, User
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from chainlit.element import Element, ElementDict
|
||||
from chainlit.step import StepDict
|
||||
|
||||
|
||||
class SQLAlchemyDataLayer(BaseDataLayer):
|
||||
def __init__(
|
||||
self,
|
||||
conninfo: str,
|
||||
connect_args: Optional[dict[str, Any]] = None,
|
||||
ssl_require: bool = False,
|
||||
storage_provider: Optional[BaseStorageClient] = None,
|
||||
user_thread_limit: Optional[int] = 1000,
|
||||
show_logger: Optional[bool] = False,
|
||||
):
|
||||
self._conninfo = conninfo
|
||||
self.user_thread_limit = user_thread_limit
|
||||
self.show_logger = show_logger
|
||||
if connect_args is None:
|
||||
connect_args = {}
|
||||
if ssl_require:
|
||||
# Create an SSL context to require an SSL connection
|
||||
ssl_context = ssl.create_default_context()
|
||||
ssl_context.check_hostname = False
|
||||
ssl_context.verify_mode = ssl.CERT_NONE
|
||||
connect_args["ssl"] = ssl_context
|
||||
self.engine: AsyncEngine = create_async_engine(
|
||||
self._conninfo, connect_args=connect_args
|
||||
)
|
||||
self.async_session = sessionmaker(
|
||||
bind=self.engine, expire_on_commit=False, class_=AsyncSession
|
||||
) # type: ignore
|
||||
if storage_provider:
|
||||
self.storage_provider: Optional[BaseStorageClient] = storage_provider
|
||||
if self.show_logger:
|
||||
logger.info("SQLAlchemyDataLayer storage client initialized")
|
||||
else:
|
||||
self.storage_provider = None
|
||||
logger.warning(
|
||||
"SQLAlchemyDataLayer storage client is not initialized and elements will not be persisted!"
|
||||
)
|
||||
|
||||
async def build_debug_url(self) -> str:
|
||||
return ""
|
||||
|
||||
###### SQL Helpers ######
|
||||
async def execute_sql(
|
||||
self, query: str, parameters: dict
|
||||
) -> Union[List[Dict[str, Any]], int, None]:
|
||||
parameterized_query = text(query)
|
||||
async with self.async_session() as session:
|
||||
try:
|
||||
await session.begin()
|
||||
result = await session.execute(parameterized_query, parameters)
|
||||
await session.commit()
|
||||
if result.returns_rows:
|
||||
json_result = [dict(row._mapping) for row in result.fetchall()]
|
||||
clean_json_result = self.clean_result(json_result)
|
||||
assert isinstance(clean_json_result, list) or isinstance(
|
||||
clean_json_result, int
|
||||
)
|
||||
return clean_json_result
|
||||
else:
|
||||
return result.rowcount
|
||||
except SQLAlchemyError as e:
|
||||
await session.rollback()
|
||||
logger.warning(f"An error occurred: {e}")
|
||||
return None
|
||||
except Exception as e:
|
||||
await session.rollback()
|
||||
logger.warning(f"An unexpected error occurred: {e}")
|
||||
return None
|
||||
|
||||
async def get_current_timestamp(self) -> str:
|
||||
return datetime.now().isoformat() + "Z"
|
||||
|
||||
def clean_result(self, obj):
|
||||
"""Recursively change UUID -> str and serialize dictionaries"""
|
||||
if isinstance(obj, dict):
|
||||
return {k: self.clean_result(v) for k, v in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [self.clean_result(item) for item in obj]
|
||||
elif isinstance(obj, uuid.UUID):
|
||||
return str(obj)
|
||||
return obj
|
||||
|
||||
###### User ######
|
||||
async def get_user(self, identifier: str) -> Optional[PersistedUser]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: get_user, identifier={identifier}")
|
||||
query = "SELECT * FROM users WHERE identifier = :identifier"
|
||||
parameters = {"identifier": identifier}
|
||||
result = await self.execute_sql(query=query, parameters=parameters)
|
||||
if result and isinstance(result, list):
|
||||
user_data = result[0]
|
||||
|
||||
# SQLite returns JSON as string, we most convert it. (#1137)
|
||||
metadata = user_data.get("metadata", {})
|
||||
if isinstance(metadata, str):
|
||||
metadata = json.loads(metadata)
|
||||
|
||||
assert isinstance(metadata, dict)
|
||||
assert isinstance(user_data["id"], str)
|
||||
assert isinstance(user_data["identifier"], str)
|
||||
assert isinstance(user_data["createdAt"], str)
|
||||
|
||||
return PersistedUser(
|
||||
id=user_data["id"],
|
||||
identifier=user_data["identifier"],
|
||||
createdAt=user_data["createdAt"],
|
||||
metadata=metadata,
|
||||
)
|
||||
return None
|
||||
|
||||
async def _get_user_identifer_by_id(self, user_id: str) -> str:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: _get_user_identifer_by_id, user_id={user_id}")
|
||||
query = "SELECT identifier FROM users WHERE id = :user_id"
|
||||
parameters = {"user_id": user_id}
|
||||
result = await self.execute_sql(query=query, parameters=parameters)
|
||||
|
||||
assert result
|
||||
assert isinstance(result, list)
|
||||
|
||||
return result[0]["identifier"]
|
||||
|
||||
async def _get_user_id_by_thread(self, thread_id: str) -> Optional[str]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: _get_user_id_by_thread, thread_id={thread_id}")
|
||||
query = """SELECT "userId" FROM threads WHERE id = :thread_id"""
|
||||
parameters = {"thread_id": thread_id}
|
||||
result = await self.execute_sql(query=query, parameters=parameters)
|
||||
if result:
|
||||
assert isinstance(result, list)
|
||||
return result[0]["userId"]
|
||||
|
||||
return None
|
||||
|
||||
async def create_user(self, user: User) -> Optional[PersistedUser]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: create_user, user_identifier={user.identifier}")
|
||||
existing_user: Optional[PersistedUser] = await self.get_user(user.identifier)
|
||||
user_dict: Dict[str, Any] = {
|
||||
"identifier": str(user.identifier),
|
||||
"metadata": json.dumps(user.metadata) or {},
|
||||
}
|
||||
if not existing_user: # create the user
|
||||
if self.show_logger:
|
||||
logger.info("SQLAlchemy: create_user, creating the user")
|
||||
user_dict["id"] = str(uuid.uuid4())
|
||||
user_dict["createdAt"] = await self.get_current_timestamp()
|
||||
query = """INSERT INTO users ("id", "identifier", "createdAt", "metadata") VALUES (:id, :identifier, :createdAt, :metadata)"""
|
||||
await self.execute_sql(query=query, parameters=user_dict)
|
||||
else: # update the user
|
||||
if self.show_logger:
|
||||
logger.info("SQLAlchemy: update user metadata")
|
||||
query = """UPDATE users SET "metadata" = :metadata WHERE "identifier" = :identifier"""
|
||||
await self.execute_sql(
|
||||
query=query, parameters=user_dict
|
||||
) # We want to update the metadata
|
||||
return await self.get_user(user.identifier)
|
||||
|
||||
###### Threads ######
|
||||
async def get_thread_author(self, thread_id: str) -> str:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: get_thread_author, thread_id={thread_id}")
|
||||
query = """SELECT "userIdentifier" FROM threads WHERE "id" = :id"""
|
||||
parameters = {"id": thread_id}
|
||||
result = await self.execute_sql(query=query, parameters=parameters)
|
||||
if isinstance(result, list) and result:
|
||||
author_identifier = result[0].get("userIdentifier")
|
||||
if author_identifier is not None:
|
||||
return author_identifier
|
||||
raise ValueError(f"Author not found for thread_id {thread_id}")
|
||||
|
||||
async def get_thread(self, thread_id: str) -> Optional[ThreadDict]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: get_thread, thread_id={thread_id}")
|
||||
user_threads: Optional[List[ThreadDict]] = await self.get_all_user_threads(
|
||||
thread_id=thread_id
|
||||
)
|
||||
if user_threads:
|
||||
return user_threads[0]
|
||||
else:
|
||||
return None
|
||||
|
||||
async def update_thread(
|
||||
self,
|
||||
thread_id: str,
|
||||
name: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
metadata: Optional[Dict] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: update_thread, thread_id={thread_id}")
|
||||
|
||||
user_identifier = None
|
||||
if user_id:
|
||||
user_identifier = await self._get_user_identifer_by_id(user_id)
|
||||
|
||||
has_updates = (
|
||||
metadata is not None
|
||||
or name is not None
|
||||
or user_id is not None
|
||||
or tags is not None
|
||||
)
|
||||
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
|
||||
existing = await self.execute_sql(
|
||||
query='SELECT "metadata" FROM threads WHERE "id" = :id',
|
||||
parameters={"id": thread_id},
|
||||
)
|
||||
|
||||
thread_exists = isinstance(existing, list) and len(existing) > 0
|
||||
if thread_exists and not has_updates:
|
||||
return
|
||||
|
||||
base = {}
|
||||
if isinstance(existing, list) and existing:
|
||||
raw = existing[0].get("metadata") or {}
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
base = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
base = {}
|
||||
elif isinstance(raw, dict):
|
||||
base = raw
|
||||
to_delete = {k for k, v in metadata.items() if v is None}
|
||||
incoming = {k: v for k, v in metadata.items() if v is not None}
|
||||
base = {k: v for k, v in base.items() if k not in to_delete}
|
||||
metadata = {**base, **incoming}
|
||||
|
||||
name_value = name
|
||||
if name_value is None and metadata:
|
||||
name_value = metadata.get("name")
|
||||
|
||||
is_new_thread = not thread_exists
|
||||
created_at_value = await self.get_current_timestamp() if is_new_thread else None
|
||||
|
||||
data = {
|
||||
"id": thread_id,
|
||||
"createdAt": created_at_value,
|
||||
"name": name_value,
|
||||
"userId": user_id,
|
||||
"userIdentifier": user_identifier,
|
||||
"tags": tags,
|
||||
"metadata": json.dumps(metadata),
|
||||
}
|
||||
parameters = {
|
||||
key: value for key, value in data.items() if value is not None
|
||||
} # Remove keys with None values
|
||||
columns = ", ".join(f'"{key}"' for key in parameters.keys())
|
||||
values = ", ".join(f":{key}" for key in parameters.keys())
|
||||
updates = ", ".join(
|
||||
f'"{key}" = EXCLUDED."{key}"' for key in parameters.keys() if key != "id"
|
||||
)
|
||||
query = f"""
|
||||
INSERT INTO threads ({columns})
|
||||
VALUES ({values})
|
||||
ON CONFLICT ("id") DO UPDATE
|
||||
SET {updates};
|
||||
"""
|
||||
await self.execute_sql(query=query, parameters=parameters)
|
||||
|
||||
async def delete_thread(self, thread_id: str):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: delete_thread, thread_id={thread_id}")
|
||||
|
||||
elements_query = """SELECT * FROM elements WHERE "threadId" = :id"""
|
||||
elements = await self.execute_sql(elements_query, {"id": thread_id})
|
||||
|
||||
if self.storage_provider is not None and isinstance(elements, list):
|
||||
for elem in filter(lambda x: x["objectKey"], elements):
|
||||
await self.storage_provider.delete_file(object_key=elem["objectKey"])
|
||||
|
||||
# Delete feedbacks/elements/steps/thread
|
||||
feedbacks_query = """DELETE FROM feedbacks WHERE "forId" IN (SELECT "id" FROM steps WHERE "threadId" = :id)"""
|
||||
elements_query = """DELETE FROM elements WHERE "threadId" = :id"""
|
||||
steps_query = """DELETE FROM steps WHERE "threadId" = :id"""
|
||||
thread_query = """DELETE FROM threads WHERE "id" = :id"""
|
||||
parameters = {"id": thread_id}
|
||||
await self.execute_sql(query=feedbacks_query, parameters=parameters)
|
||||
await self.execute_sql(query=elements_query, parameters=parameters)
|
||||
await self.execute_sql(query=steps_query, parameters=parameters)
|
||||
await self.execute_sql(query=thread_query, parameters=parameters)
|
||||
|
||||
async def list_threads(
|
||||
self, pagination: Pagination, filters: ThreadFilter
|
||||
) -> PaginatedResponse:
|
||||
if self.show_logger:
|
||||
logger.info(
|
||||
f"SQLAlchemy: list_threads, pagination={pagination}, filters={filters}"
|
||||
)
|
||||
if not filters.userId:
|
||||
raise ValueError("userId is required")
|
||||
all_user_threads: List[ThreadDict] = (
|
||||
await self.get_all_user_threads(user_id=filters.userId) or []
|
||||
)
|
||||
|
||||
search_keyword = filters.search.lower() if filters.search else None
|
||||
feedback_value = int(filters.feedback) if filters.feedback else None
|
||||
|
||||
filtered_threads = []
|
||||
for thread in all_user_threads:
|
||||
keyword_match = True
|
||||
feedback_match = True
|
||||
if search_keyword or feedback_value is not None:
|
||||
if search_keyword:
|
||||
keyword_match = any(
|
||||
search_keyword in step["output"].lower()
|
||||
for step in thread["steps"]
|
||||
if "output" in step
|
||||
)
|
||||
if feedback_value is not None:
|
||||
feedback_match = False # Assume no match until found
|
||||
for step in thread["steps"]:
|
||||
feedback = step.get("feedback")
|
||||
if feedback and feedback.get("value") == feedback_value:
|
||||
feedback_match = True
|
||||
break
|
||||
if keyword_match and feedback_match:
|
||||
filtered_threads.append(thread)
|
||||
|
||||
start = 0
|
||||
if pagination.cursor:
|
||||
for i, thread in enumerate(filtered_threads):
|
||||
if (
|
||||
thread["id"] == pagination.cursor
|
||||
): # Find the start index using pagination.cursor
|
||||
start = i + 1
|
||||
break
|
||||
end = start + pagination.first
|
||||
paginated_threads = filtered_threads[start:end] or []
|
||||
|
||||
has_next_page = len(filtered_threads) > end
|
||||
start_cursor = paginated_threads[0]["id"] if paginated_threads else None
|
||||
end_cursor = paginated_threads[-1]["id"] if paginated_threads else None
|
||||
|
||||
return PaginatedResponse(
|
||||
pageInfo=PageInfo(
|
||||
hasNextPage=has_next_page,
|
||||
startCursor=start_cursor,
|
||||
endCursor=end_cursor,
|
||||
),
|
||||
data=paginated_threads,
|
||||
)
|
||||
|
||||
###### Steps ######
|
||||
@queue_until_user_message()
|
||||
async def create_step(self, step_dict: "StepDict"):
|
||||
await self.update_thread(step_dict["threadId"])
|
||||
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: create_step, step_id={step_dict.get('id')}")
|
||||
|
||||
step_dict["showInput"] = (
|
||||
str(step_dict.get("showInput", "")).lower()
|
||||
if "showInput" in step_dict
|
||||
else None
|
||||
)
|
||||
parameters = {
|
||||
key: value
|
||||
for key, value in step_dict.items()
|
||||
if value is not None and not (isinstance(value, dict) and not value)
|
||||
}
|
||||
parameters["metadata"] = json.dumps(step_dict.get("metadata", {}))
|
||||
parameters["generation"] = json.dumps(step_dict.get("generation", {}))
|
||||
columns = ", ".join(f'"{key}"' for key in parameters.keys())
|
||||
values = ", ".join(f":{key}" for key in parameters.keys())
|
||||
updates = ", ".join(
|
||||
f'"{key}" = :{key}' for key in parameters.keys() if key != "id"
|
||||
)
|
||||
query = f"""
|
||||
INSERT INTO steps ({columns})
|
||||
VALUES ({values})
|
||||
ON CONFLICT (id) DO UPDATE
|
||||
SET {updates};
|
||||
"""
|
||||
await self.execute_sql(query=query, parameters=parameters)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def update_step(self, step_dict: "StepDict"):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: update_step, step_id={step_dict.get('id')}")
|
||||
await self.create_step(step_dict)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_step(self, step_id: str):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: delete_step, step_id={step_id}")
|
||||
# Delete feedbacks/elements/steps
|
||||
feedbacks_query = """DELETE FROM feedbacks WHERE "forId" = :id"""
|
||||
elements_query = """DELETE FROM elements WHERE "forId" = :id"""
|
||||
steps_query = """DELETE FROM steps WHERE "id" = :id"""
|
||||
parameters = {"id": step_id}
|
||||
await self.execute_sql(query=feedbacks_query, parameters=parameters)
|
||||
await self.execute_sql(query=elements_query, parameters=parameters)
|
||||
await self.execute_sql(query=steps_query, parameters=parameters)
|
||||
|
||||
async def get_step(self, step_id: str) -> Optional["StepDict"]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: get_step, step_id={step_id}")
|
||||
steps_feedbacks_query = """
|
||||
SELECT
|
||||
s."id" AS step_id,
|
||||
s."name" AS step_name,
|
||||
s."type" AS step_type,
|
||||
s."threadId" AS step_threadid,
|
||||
s."parentId" AS step_parentid,
|
||||
s."streaming" AS step_streaming,
|
||||
s."waitForAnswer" AS step_waitforanswer,
|
||||
s."isError" AS step_iserror,
|
||||
s."metadata" AS step_metadata,
|
||||
s."tags" AS step_tags,
|
||||
s."input" AS step_input,
|
||||
s."output" AS step_output,
|
||||
s."createdAt" AS step_createdat,
|
||||
s."start" AS step_start,
|
||||
s."end" AS step_end,
|
||||
s."generation" AS step_generation,
|
||||
s."showInput" AS step_showinput,
|
||||
s."language" AS step_language,
|
||||
f."value" AS feedback_value,
|
||||
f."comment" AS feedback_comment,
|
||||
f."id" AS feedback_id
|
||||
FROM steps s LEFT JOIN feedbacks f ON s."id" = f."forId"
|
||||
WHERE s."id" = :step_id
|
||||
"""
|
||||
steps_feedbacks = await self.execute_sql(
|
||||
query=steps_feedbacks_query, parameters={"step_id": step_id}
|
||||
)
|
||||
|
||||
if not isinstance(steps_feedbacks, list) or not steps_feedbacks:
|
||||
return None
|
||||
|
||||
step_feedback = steps_feedbacks[0]
|
||||
|
||||
feedback = None
|
||||
if step_feedback["feedback_value"] is not None:
|
||||
feedback = FeedbackDict(
|
||||
forId=step_feedback["step_id"],
|
||||
id=step_feedback.get("feedback_id"),
|
||||
value=step_feedback["feedback_value"],
|
||||
comment=step_feedback.get("feedback_comment"),
|
||||
)
|
||||
return StepDict(
|
||||
id=step_feedback["step_id"],
|
||||
name=step_feedback["step_name"],
|
||||
type=step_feedback["step_type"],
|
||||
threadId=step_feedback.get("step_threadid", ""),
|
||||
parentId=step_feedback.get("step_parentid"),
|
||||
streaming=step_feedback.get("step_streaming", False),
|
||||
waitForAnswer=step_feedback.get("step_waitforanswer"),
|
||||
isError=step_feedback.get("step_iserror"),
|
||||
metadata=(
|
||||
step_feedback["step_metadata"]
|
||||
if step_feedback.get("step_metadata") is not None
|
||||
else {}
|
||||
),
|
||||
tags=step_feedback.get("step_tags"),
|
||||
input=(
|
||||
step_feedback.get("step_input", "")
|
||||
if step_feedback.get("step_showinput") not in [None, "false"]
|
||||
else ""
|
||||
),
|
||||
output=step_feedback.get("step_output", ""),
|
||||
createdAt=step_feedback.get("step_createdat"),
|
||||
start=step_feedback.get("step_start"),
|
||||
end=step_feedback.get("step_end"),
|
||||
generation=step_feedback.get("step_generation"),
|
||||
showInput=step_feedback.get("step_showinput"),
|
||||
language=step_feedback.get("step_language"),
|
||||
feedback=feedback,
|
||||
)
|
||||
|
||||
###### Feedback ######
|
||||
async def upsert_feedback(self, feedback: Feedback) -> str:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: upsert_feedback, feedback_id={feedback.id}")
|
||||
feedback.id = feedback.id or str(uuid.uuid4())
|
||||
feedback_dict = asdict(feedback)
|
||||
parameters = {
|
||||
key: value for key, value in feedback_dict.items() if value is not None
|
||||
}
|
||||
|
||||
columns = ", ".join(f'"{key}"' for key in parameters.keys())
|
||||
values = ", ".join(f":{key}" for key in parameters.keys())
|
||||
updates = ", ".join(
|
||||
f'"{key}" = :{key}' for key in parameters.keys() if key != "id"
|
||||
)
|
||||
query = f"""
|
||||
INSERT INTO feedbacks ({columns})
|
||||
VALUES ({values})
|
||||
ON CONFLICT (id) DO UPDATE
|
||||
SET {updates};
|
||||
"""
|
||||
await self.execute_sql(query=query, parameters=parameters)
|
||||
return feedback.id
|
||||
|
||||
async def delete_feedback(self, feedback_id: str) -> bool:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: delete_feedback, feedback_id={feedback_id}")
|
||||
query = """DELETE FROM feedbacks WHERE "id" = :feedback_id"""
|
||||
parameters = {"feedback_id": feedback_id}
|
||||
await self.execute_sql(query=query, parameters=parameters)
|
||||
return True
|
||||
|
||||
###### Elements ######
|
||||
async def get_element(
|
||||
self, thread_id: str, element_id: str
|
||||
) -> Optional["ElementDict"]:
|
||||
if self.show_logger:
|
||||
logger.info(
|
||||
f"SQLAlchemy: get_element, thread_id={thread_id}, element_id={element_id}"
|
||||
)
|
||||
query = """SELECT * FROM elements WHERE "threadId" = :thread_id AND "id" = :element_id"""
|
||||
parameters = {"thread_id": thread_id, "element_id": element_id}
|
||||
element: Union[List[Dict[str, Any]], int, None] = await self.execute_sql(
|
||||
query=query, parameters=parameters
|
||||
)
|
||||
if isinstance(element, list) and element:
|
||||
element_dict: Dict[str, Any] = element[0]
|
||||
return ElementDict(
|
||||
id=element_dict["id"],
|
||||
threadId=element_dict.get("threadId"),
|
||||
type=element_dict["type"],
|
||||
chainlitKey=element_dict.get("chainlitKey"),
|
||||
url=element_dict.get("url"),
|
||||
objectKey=element_dict.get("objectKey"),
|
||||
name=element_dict["name"],
|
||||
props=json.loads(element_dict.get("props", "{}")),
|
||||
display=element_dict["display"],
|
||||
size=element_dict.get("size"),
|
||||
language=element_dict.get("language"),
|
||||
page=element_dict.get("page"),
|
||||
autoPlay=element_dict.get("autoPlay"),
|
||||
playerConfig=element_dict.get("playerConfig"),
|
||||
forId=element_dict.get("forId"),
|
||||
mime=element_dict.get("mime"),
|
||||
)
|
||||
else:
|
||||
return None
|
||||
|
||||
@queue_until_user_message()
|
||||
async def create_element(self, element: "Element"):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: create_element, element_id = {element.id}")
|
||||
|
||||
if not self.storage_provider:
|
||||
logger.warning(
|
||||
"SQLAlchemy: create_element error. No blob_storage_client is configured!"
|
||||
)
|
||||
return
|
||||
if not element.for_id:
|
||||
return
|
||||
|
||||
content: Optional[Union[bytes, str]] = None
|
||||
|
||||
if element.path:
|
||||
async with aiofiles.open(element.path, "rb") as f:
|
||||
content = await f.read()
|
||||
elif element.url:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(element.url) as response:
|
||||
if response.status == 200:
|
||||
content = await response.read()
|
||||
else:
|
||||
content = None
|
||||
elif element.content:
|
||||
content = element.content
|
||||
else:
|
||||
raise ValueError("Element url, path or content must be provided")
|
||||
if content is None:
|
||||
raise ValueError("Content is None, cannot upload file")
|
||||
|
||||
user_id: str = await self._get_user_id_by_thread(element.thread_id) or "unknown"
|
||||
file_object_key = f"{user_id}/{element.id}" + (
|
||||
f"/{element.name}" if element.name else ""
|
||||
)
|
||||
|
||||
if not element.mime:
|
||||
element.mime = "application/octet-stream"
|
||||
|
||||
uploaded_file = await self.storage_provider.upload_file(
|
||||
object_key=file_object_key, data=content, mime=element.mime, overwrite=True
|
||||
)
|
||||
if not uploaded_file:
|
||||
raise ValueError(
|
||||
"SQLAlchemy Error: create_element, Failed to persist data in storage_provider"
|
||||
)
|
||||
|
||||
element_dict: ElementDict = element.to_dict()
|
||||
|
||||
element_dict["url"] = uploaded_file.get("url")
|
||||
element_dict["objectKey"] = uploaded_file.get("object_key")
|
||||
|
||||
element_dict_cleaned = {k: v for k, v in element_dict.items() if v is not None}
|
||||
if "props" in element_dict_cleaned:
|
||||
element_dict_cleaned["props"] = json.dumps(element_dict_cleaned["props"])
|
||||
|
||||
columns = ", ".join(f'"{column}"' for column in element_dict_cleaned.keys())
|
||||
placeholders = ", ".join(f":{column}" for column in element_dict_cleaned.keys())
|
||||
updates = ", ".join(
|
||||
f'"{column}" = :{column}'
|
||||
for column in element_dict_cleaned.keys()
|
||||
if column != "id"
|
||||
)
|
||||
query = f"INSERT INTO elements ({columns}) VALUES ({placeholders}) ON CONFLICT (id) DO UPDATE SET {updates};"
|
||||
await self.execute_sql(query=query, parameters=element_dict_cleaned)
|
||||
|
||||
@queue_until_user_message()
|
||||
async def delete_element(self, element_id: str, thread_id: Optional[str] = None):
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: delete_element, element_id={element_id}")
|
||||
|
||||
query = """SELECT * FROM elements WHERE "id" = :id"""
|
||||
elements = await self.execute_sql(query, {"id": element_id})
|
||||
|
||||
if (
|
||||
self.storage_provider is not None
|
||||
and isinstance(elements, list)
|
||||
and len(elements) > 0
|
||||
and elements[0]["objectKey"]
|
||||
):
|
||||
await self.storage_provider.delete_file(object_key=elements[0]["objectKey"])
|
||||
|
||||
query = """DELETE FROM elements WHERE "id" = :id"""
|
||||
parameters = {"id": element_id}
|
||||
|
||||
await self.execute_sql(query=query, parameters=parameters)
|
||||
|
||||
async def get_all_user_threads(
|
||||
self, user_id: Optional[str] = None, thread_id: Optional[str] = None
|
||||
) -> Optional[List[ThreadDict]]:
|
||||
"""Fetch all user threads up to self.user_thread_limit, or one thread by id if thread_id is provided."""
|
||||
if self.show_logger:
|
||||
logger.info("SQLAlchemy: get_all_user_threads")
|
||||
user_threads_query = """
|
||||
SELECT
|
||||
t."id" AS thread_id,
|
||||
t."createdAt" AS thread_createdat,
|
||||
t."name" AS thread_name,
|
||||
t."userId" AS user_id,
|
||||
t."userIdentifier" AS user_identifier,
|
||||
t."tags" AS thread_tags,
|
||||
t."metadata" AS thread_metadata,
|
||||
MAX(s."createdAt") AS updatedAt
|
||||
FROM threads t
|
||||
LEFT JOIN steps s ON t."id" = s."threadId"
|
||||
WHERE t."userId" = :user_id OR t."id" = :thread_id
|
||||
GROUP BY
|
||||
t."id",
|
||||
t."createdAt",
|
||||
t."name",
|
||||
t."userId",
|
||||
t."userIdentifier",
|
||||
t."tags",
|
||||
t."metadata"
|
||||
ORDER BY updatedAt DESC NULLS LAST
|
||||
LIMIT :limit
|
||||
"""
|
||||
user_threads = await self.execute_sql(
|
||||
query=user_threads_query,
|
||||
parameters={
|
||||
"user_id": user_id,
|
||||
"limit": self.user_thread_limit,
|
||||
"thread_id": thread_id,
|
||||
},
|
||||
)
|
||||
if not isinstance(user_threads, list):
|
||||
return None
|
||||
if not user_threads:
|
||||
return []
|
||||
else:
|
||||
thread_ids = (
|
||||
"('"
|
||||
+ "','".join(map(str, [thread["thread_id"] for thread in user_threads]))
|
||||
+ "')"
|
||||
)
|
||||
|
||||
steps_feedbacks_query = f"""
|
||||
SELECT
|
||||
s."id" AS step_id,
|
||||
s."name" AS step_name,
|
||||
s."type" AS step_type,
|
||||
s."threadId" AS step_threadid,
|
||||
s."parentId" AS step_parentid,
|
||||
s."streaming" AS step_streaming,
|
||||
s."waitForAnswer" AS step_waitforanswer,
|
||||
s."isError" AS step_iserror,
|
||||
s."metadata" AS step_metadata,
|
||||
s."tags" AS step_tags,
|
||||
s."input" AS step_input,
|
||||
s."output" AS step_output,
|
||||
s."createdAt" AS step_createdat,
|
||||
s."start" AS step_start,
|
||||
s."end" AS step_end,
|
||||
s."generation" AS step_generation,
|
||||
s."showInput" AS step_showinput,
|
||||
s."language" AS step_language,
|
||||
f."value" AS feedback_value,
|
||||
f."comment" AS feedback_comment,
|
||||
f."id" AS feedback_id
|
||||
FROM steps s LEFT JOIN feedbacks f ON s."id" = f."forId"
|
||||
WHERE s."threadId" IN {thread_ids}
|
||||
ORDER BY s."createdAt" ASC
|
||||
"""
|
||||
steps_feedbacks = await self.execute_sql(
|
||||
query=steps_feedbacks_query, parameters={}
|
||||
)
|
||||
|
||||
elements_query = f"""
|
||||
SELECT
|
||||
e."id" AS element_id,
|
||||
e."threadId" as element_threadid,
|
||||
e."type" AS element_type,
|
||||
e."chainlitKey" AS element_chainlitkey,
|
||||
e."url" AS element_url,
|
||||
e."objectKey" as element_objectkey,
|
||||
e."name" AS element_name,
|
||||
e."display" AS element_display,
|
||||
e."size" AS element_size,
|
||||
e."language" AS element_language,
|
||||
e."page" AS element_page,
|
||||
e."forId" AS element_forid,
|
||||
e."mime" AS element_mime,
|
||||
e."props" AS props
|
||||
FROM elements e
|
||||
WHERE e."threadId" IN {thread_ids}
|
||||
"""
|
||||
elements = await self.execute_sql(query=elements_query, parameters={})
|
||||
|
||||
thread_dicts = {}
|
||||
for thread in user_threads:
|
||||
thread_id = thread["thread_id"]
|
||||
if thread_id is not None:
|
||||
thread_dicts[thread_id] = ThreadDict(
|
||||
id=thread_id,
|
||||
createdAt=thread["thread_createdat"],
|
||||
name=thread["thread_name"],
|
||||
userId=thread["user_id"],
|
||||
userIdentifier=thread["user_identifier"],
|
||||
tags=thread["thread_tags"],
|
||||
metadata=thread["thread_metadata"],
|
||||
steps=[],
|
||||
elements=[],
|
||||
)
|
||||
# Process steps_feedbacks to populate the steps in the corresponding ThreadDict
|
||||
if isinstance(steps_feedbacks, list):
|
||||
for step_feedback in steps_feedbacks:
|
||||
thread_id = step_feedback["step_threadid"]
|
||||
if thread_id is not None:
|
||||
feedback = None
|
||||
if step_feedback["feedback_value"] is not None:
|
||||
feedback = FeedbackDict(
|
||||
forId=step_feedback["step_id"],
|
||||
id=step_feedback.get("feedback_id"),
|
||||
value=step_feedback["feedback_value"],
|
||||
comment=step_feedback.get("feedback_comment"),
|
||||
)
|
||||
step_dict = StepDict(
|
||||
id=step_feedback["step_id"],
|
||||
name=step_feedback["step_name"],
|
||||
type=step_feedback["step_type"],
|
||||
threadId=thread_id,
|
||||
parentId=step_feedback.get("step_parentid"),
|
||||
streaming=step_feedback.get("step_streaming", False),
|
||||
waitForAnswer=step_feedback.get("step_waitforanswer"),
|
||||
isError=step_feedback.get("step_iserror"),
|
||||
metadata=(
|
||||
step_feedback["step_metadata"]
|
||||
if step_feedback.get("step_metadata") is not None
|
||||
else {}
|
||||
),
|
||||
tags=step_feedback.get("step_tags"),
|
||||
input=(
|
||||
step_feedback.get("step_input", "")
|
||||
if step_feedback.get("step_showinput")
|
||||
not in [None, "false"]
|
||||
else ""
|
||||
),
|
||||
output=step_feedback.get("step_output", ""),
|
||||
createdAt=step_feedback.get("step_createdat"),
|
||||
start=step_feedback.get("step_start"),
|
||||
end=step_feedback.get("step_end"),
|
||||
generation=step_feedback.get("step_generation"),
|
||||
showInput=step_feedback.get("step_showinput"),
|
||||
language=step_feedback.get("step_language"),
|
||||
feedback=feedback,
|
||||
)
|
||||
# Append the step to the steps list of the corresponding ThreadDict
|
||||
thread_dicts[thread_id]["steps"].append(step_dict)
|
||||
|
||||
if isinstance(elements, list):
|
||||
for element in elements:
|
||||
thread_id = element["element_threadid"]
|
||||
if thread_id is not None:
|
||||
element_url: str | None = None
|
||||
object_key_val = element.get("element_objectkey")
|
||||
if (
|
||||
self.storage_provider is not None
|
||||
and isinstance(object_key_val, str)
|
||||
and object_key_val.strip()
|
||||
):
|
||||
try:
|
||||
element_url = await self.storage_provider.get_read_url(
|
||||
object_key=object_key_val,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to get read URL for object_key '{object_key_val}': {e}. Falling back to stored URL."
|
||||
)
|
||||
element_url = element.get("element_url")
|
||||
else:
|
||||
element_url = element.get("element_url")
|
||||
element_dict = ElementDict(
|
||||
id=element["element_id"],
|
||||
threadId=thread_id,
|
||||
type=element["element_type"],
|
||||
chainlitKey=element.get("element_chainlitkey"),
|
||||
url=element_url,
|
||||
objectKey=element.get("element_objectkey"),
|
||||
name=element["element_name"],
|
||||
display=element["element_display"],
|
||||
size=element.get("element_size"),
|
||||
language=element.get("element_language"),
|
||||
autoPlay=element.get("element_autoPlay"),
|
||||
playerConfig=element.get("element_playerconfig"),
|
||||
page=element.get("element_page"),
|
||||
props=element.get("props", "{}"),
|
||||
forId=element.get("element_forid"),
|
||||
mime=element.get("element_mime"),
|
||||
)
|
||||
thread_dicts[thread_id]["elements"].append(element_dict) # type: ignore
|
||||
|
||||
return list(thread_dicts.values())
|
||||
|
||||
async def get_favorite_steps(self, user_id: str) -> List[StepDict]:
|
||||
if self.show_logger:
|
||||
logger.info(f"SQLAlchemy: get_favorite_steps, user_id={user_id}")
|
||||
|
||||
query = """
|
||||
SELECT
|
||||
s."id" AS step_id,
|
||||
s."name" AS step_name,
|
||||
s."type" AS step_type,
|
||||
s."threadId" AS step_threadid,
|
||||
s."parentId" AS step_parentid,
|
||||
s."streaming" AS step_streaming,
|
||||
s."waitForAnswer" AS step_waitforanswer,
|
||||
s."isError" AS step_iserror,
|
||||
s."metadata" AS step_metadata,
|
||||
s."tags" AS step_tags,
|
||||
s."input" AS step_input,
|
||||
s."output" AS step_output,
|
||||
s."createdAt" AS step_createdat,
|
||||
s."start" AS step_start,
|
||||
s."end" AS step_end,
|
||||
s."generation" AS step_generation,
|
||||
s."showInput" AS step_showinput,
|
||||
s."language" AS step_language
|
||||
FROM steps s
|
||||
JOIN threads t ON s."threadId" = t.id
|
||||
WHERE t."userId" = :user_id
|
||||
AND s."metadata" LIKE :favorite_pattern
|
||||
ORDER BY s."createdAt" DESC \
|
||||
"""
|
||||
|
||||
result = await self.execute_sql(
|
||||
query, {"user_id": user_id, "favorite_pattern": '%"favorite": true%'}
|
||||
)
|
||||
|
||||
steps = []
|
||||
if isinstance(result, list):
|
||||
for row in result:
|
||||
metadata_raw = row["step_metadata"]
|
||||
meta_dict = {}
|
||||
if isinstance(metadata_raw, str):
|
||||
try:
|
||||
meta_dict = json.loads(metadata_raw)
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(metadata_raw, dict):
|
||||
meta_dict = metadata_raw
|
||||
|
||||
if meta_dict.get("favorite"):
|
||||
steps.append(
|
||||
StepDict(
|
||||
id=row["step_id"],
|
||||
name=row["step_name"],
|
||||
type=row["step_type"],
|
||||
threadId=row["step_threadid"],
|
||||
parentId=row["step_parentid"],
|
||||
streaming=row.get("step_streaming", False),
|
||||
waitForAnswer=row.get("step_waitforanswer"),
|
||||
isError=row.get("step_iserror"),
|
||||
metadata=meta_dict,
|
||||
tags=row.get("step_tags"),
|
||||
input=(
|
||||
row.get("step_input", "")
|
||||
if row.get("step_showinput") not in [None, "false"]
|
||||
else ""
|
||||
),
|
||||
output=row.get("step_output", ""),
|
||||
createdAt=row.get("step_createdat"),
|
||||
start=row.get("step_start"),
|
||||
end=row.get("step_end"),
|
||||
generation=row.get("step_generation"),
|
||||
showInput=row.get("step_showinput"),
|
||||
language=row.get("step_language"),
|
||||
feedback=None,
|
||||
)
|
||||
)
|
||||
return steps
|
||||
|
||||
async def close(self) -> None:
|
||||
if self.storage_provider:
|
||||
await self.storage_provider.close()
|
||||
await self.engine.dispose()
|
||||
@@ -0,0 +1,88 @@
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
|
||||
|
||||
from azure.storage.filedatalake import (
|
||||
ContentSettings,
|
||||
DataLakeFileClient,
|
||||
DataLakeServiceClient,
|
||||
FileSystemClient,
|
||||
)
|
||||
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient
|
||||
from chainlit.logger import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from azure.core.credentials import (
|
||||
AzureNamedKeyCredential,
|
||||
AzureSasCredential,
|
||||
TokenCredential,
|
||||
)
|
||||
|
||||
|
||||
class AzureStorageClient(BaseStorageClient):
|
||||
"""
|
||||
Class to enable Azure Data Lake Storage (ADLS) Gen2
|
||||
|
||||
parms:
|
||||
account_url: "https://<your_account>.dfs.core.windows.net"
|
||||
credential: Access credential (AzureKeyCredential)
|
||||
sas_token: Optionally include SAS token to append to urls
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
account_url: str,
|
||||
container: str,
|
||||
credential: Optional[
|
||||
Union[
|
||||
str,
|
||||
Dict[str, str],
|
||||
"AzureNamedKeyCredential",
|
||||
"AzureSasCredential",
|
||||
"TokenCredential",
|
||||
]
|
||||
],
|
||||
sas_token: Optional[str] = None,
|
||||
):
|
||||
try:
|
||||
self.data_lake_client = DataLakeServiceClient(
|
||||
account_url=account_url, credential=credential
|
||||
)
|
||||
self.container_client: FileSystemClient = (
|
||||
self.data_lake_client.get_file_system_client(file_system=container)
|
||||
)
|
||||
self.sas_token = sas_token
|
||||
logger.info("AzureStorageClient initialized")
|
||||
except Exception as e:
|
||||
logger.warning(f"AzureStorageClient initialization error: {e}")
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
file_client: DataLakeFileClient = self.container_client.get_file_client(
|
||||
object_key
|
||||
)
|
||||
content_settings = ContentSettings(
|
||||
content_type=mime, content_disposition=content_disposition
|
||||
)
|
||||
file_client.upload_data(
|
||||
data, overwrite=overwrite, content_settings=content_settings
|
||||
)
|
||||
url = (
|
||||
f"{file_client.url}{self.sas_token}"
|
||||
if self.sas_token
|
||||
else file_client.url
|
||||
)
|
||||
return {"object_key": object_key, "url": url}
|
||||
except Exception as e:
|
||||
logger.warning(f"AzureStorageClient, upload_file error: {e}")
|
||||
return {}
|
||||
|
||||
async def close(self) -> None:
|
||||
self.container_client.close()
|
||||
self.data_lake_client.close()
|
||||
@@ -0,0 +1,98 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, Union
|
||||
|
||||
from azure.storage.blob import BlobSasPermissions, ContentSettings, generate_blob_sas
|
||||
from azure.storage.blob.aio import BlobServiceClient as AsyncBlobServiceClient
|
||||
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient, storage_expiry_time
|
||||
from chainlit.logger import logger
|
||||
|
||||
|
||||
class AzureBlobStorageClient(BaseStorageClient):
|
||||
def __init__(self, container_name: str, storage_account: str, storage_key: str):
|
||||
self.container_name = container_name
|
||||
self.storage_account = storage_account
|
||||
self.storage_key = storage_key
|
||||
connection_string = (
|
||||
f"DefaultEndpointsProtocol=https;"
|
||||
f"AccountName={storage_account};"
|
||||
f"AccountKey={storage_key};"
|
||||
f"EndpointSuffix=core.windows.net"
|
||||
)
|
||||
self.service_client = AsyncBlobServiceClient.from_connection_string(
|
||||
connection_string
|
||||
)
|
||||
self.container_client = self.service_client.get_container_client(
|
||||
self.container_name
|
||||
)
|
||||
logger.info("AzureBlobStorageClient initialized")
|
||||
|
||||
async def get_read_url(self, object_key: str) -> str:
|
||||
if not self.storage_key:
|
||||
raise Exception("Not using Azure Storage")
|
||||
|
||||
sas_permissions = BlobSasPermissions(read=True)
|
||||
start_time = datetime.now(tz=timezone.utc)
|
||||
expiry_time = start_time + timedelta(seconds=storage_expiry_time)
|
||||
|
||||
sas_token = generate_blob_sas(
|
||||
account_name=self.storage_account,
|
||||
container_name=self.container_name,
|
||||
blob_name=object_key,
|
||||
account_key=self.storage_key,
|
||||
permission=sas_permissions,
|
||||
start=start_time,
|
||||
expiry=expiry_time,
|
||||
)
|
||||
|
||||
return f"https://{self.storage_account}.blob.core.windows.net/{self.container_name}/{object_key}?{sas_token}"
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
blob_client = self.container_client.get_blob_client(object_key)
|
||||
|
||||
if isinstance(data, str):
|
||||
data = data.encode("utf-8")
|
||||
|
||||
content_settings = ContentSettings(
|
||||
content_type=mime, content_disposition=content_disposition
|
||||
)
|
||||
|
||||
await blob_client.upload_blob(
|
||||
data, overwrite=overwrite, content_settings=content_settings
|
||||
)
|
||||
|
||||
properties = await blob_client.get_blob_properties()
|
||||
|
||||
return {
|
||||
"path": object_key,
|
||||
"object_key": object_key,
|
||||
"url": await self.get_read_url(object_key),
|
||||
"size": properties.size,
|
||||
"last_modified": properties.last_modified,
|
||||
"etag": properties.etag,
|
||||
"content_type": properties.content_settings.content_type,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to upload file to Azure Blob Storage: {e!s}")
|
||||
|
||||
async def delete_file(self, object_key: str) -> bool:
|
||||
try:
|
||||
blob_client = self.container_client.get_blob_client(blob=object_key)
|
||||
await blob_client.delete_blob()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"AzureBlobStorageClient, delete_file error: {e}")
|
||||
return False
|
||||
|
||||
async def close(self) -> None:
|
||||
await self.container_client.close()
|
||||
await self.service_client.close()
|
||||
@@ -0,0 +1,32 @@
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Union
|
||||
|
||||
storage_expiry_time = int(os.getenv("STORAGE_EXPIRY_TIME", 3600))
|
||||
|
||||
|
||||
class BaseStorageClient(ABC):
|
||||
"""Base class for non-text data persistence like Azure Data Lake, S3, Google Storage, etc."""
|
||||
|
||||
@abstractmethod
|
||||
async def upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def delete_file(self, object_key: str) -> bool:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_read_url(self, object_key: str) -> str:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def close(self) -> None:
|
||||
pass
|
||||
@@ -0,0 +1,104 @@
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
from google.auth import default
|
||||
from google.cloud import storage # type: ignore
|
||||
from google.oauth2 import service_account
|
||||
|
||||
from chainlit import make_async
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient, storage_expiry_time
|
||||
from chainlit.logger import logger
|
||||
|
||||
|
||||
class GCSStorageClient(BaseStorageClient):
|
||||
def __init__(
|
||||
self,
|
||||
bucket_name: str,
|
||||
project_id: Optional[str] = None,
|
||||
client_email: Optional[str] = None,
|
||||
private_key: Optional[str] = None,
|
||||
):
|
||||
if client_email and private_key and project_id:
|
||||
# Go to IAM & Admin, click on Service Accounts, and generate a new JSON key
|
||||
logger.info("Using Private Key from Environment Variable")
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
{
|
||||
"type": "service_account",
|
||||
"project_id": project_id,
|
||||
"private_key": private_key,
|
||||
"client_email": client_email,
|
||||
"token_uri": "https://oauth2.googleapis.com/token",
|
||||
}
|
||||
)
|
||||
else:
|
||||
# Application Default Credentials (e.g. in Google Cloud Run)
|
||||
logger.info("Using Application Default Credentials.")
|
||||
credentials, default_project_id = default()
|
||||
if not project_id:
|
||||
project_id = default_project_id
|
||||
|
||||
self.client = storage.Client(project=project_id, credentials=credentials)
|
||||
self.bucket = self.client.bucket(bucket_name)
|
||||
logger.info("GCSStorageClient initialized")
|
||||
|
||||
def sync_get_read_url(self, object_key: str) -> str:
|
||||
return self.bucket.blob(object_key).generate_signed_url(
|
||||
version="v4", expiration=storage_expiry_time, method="GET"
|
||||
)
|
||||
|
||||
async def get_read_url(self, object_key: str) -> str:
|
||||
return await make_async(self.sync_get_read_url)(object_key)
|
||||
|
||||
def sync_upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
blob = self.bucket.blob(object_key)
|
||||
|
||||
if not overwrite and blob.exists():
|
||||
raise Exception(
|
||||
f"File {object_key} already exists and overwrite is False"
|
||||
)
|
||||
|
||||
if isinstance(data, str):
|
||||
data = data.encode("utf-8")
|
||||
|
||||
blob.upload_from_string(data, content_type=mime)
|
||||
|
||||
# Return signed URL
|
||||
return {
|
||||
"object_key": object_key,
|
||||
"url": self.sync_get_read_url(object_key),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to upload file to GCS: {e!s}")
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
return await make_async(self.sync_upload_file)(
|
||||
object_key, data, mime, overwrite
|
||||
)
|
||||
|
||||
def sync_delete_file(self, object_key: str) -> bool:
|
||||
try:
|
||||
self.bucket.blob(object_key).delete()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"GCSStorageClient, delete_file error: {e}")
|
||||
return False
|
||||
|
||||
async def delete_file(self, object_key: str) -> bool:
|
||||
return await make_async(self.sync_delete_file)(object_key)
|
||||
|
||||
async def close(self) -> None:
|
||||
self.client.close()
|
||||
@@ -0,0 +1,91 @@
|
||||
import os
|
||||
from typing import Any, Dict, Union
|
||||
|
||||
import boto3 # type: ignore
|
||||
|
||||
from chainlit import make_async
|
||||
from chainlit.data.storage_clients.base import BaseStorageClient, storage_expiry_time
|
||||
from chainlit.logger import logger
|
||||
|
||||
|
||||
class S3StorageClient(BaseStorageClient):
|
||||
"""
|
||||
Class to enable Amazon S3 storage provider
|
||||
"""
|
||||
|
||||
def __init__(self, bucket: str, **kwargs: Any):
|
||||
try:
|
||||
self.bucket = bucket
|
||||
self.client = boto3.client("s3", **kwargs)
|
||||
logger.info("S3StorageClient initialized")
|
||||
except Exception as e:
|
||||
logger.warning(f"S3StorageClient initialization error: {e}")
|
||||
|
||||
def sync_get_read_url(self, object_key: str) -> str:
|
||||
try:
|
||||
url = self.client.generate_presigned_url(
|
||||
"get_object",
|
||||
Params={"Bucket": self.bucket, "Key": object_key},
|
||||
ExpiresIn=storage_expiry_time,
|
||||
)
|
||||
return url
|
||||
except Exception as e:
|
||||
logger.warning(f"S3StorageClient, get_read_url error: {e}")
|
||||
return object_key
|
||||
|
||||
async def get_read_url(self, object_key: str) -> str:
|
||||
return await make_async(self.sync_get_read_url)(object_key)
|
||||
|
||||
def sync_upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
if content_disposition is not None:
|
||||
self.client.put_object(
|
||||
Bucket=self.bucket,
|
||||
Key=object_key,
|
||||
Body=data,
|
||||
ContentType=mime,
|
||||
ContentDisposition=content_disposition,
|
||||
)
|
||||
else:
|
||||
self.client.put_object(
|
||||
Bucket=self.bucket, Key=object_key, Body=data, ContentType=mime
|
||||
)
|
||||
endpoint = os.environ.get("DEV_AWS_ENDPOINT", "amazonaws.com")
|
||||
url = f"https://{self.bucket}.s3.{endpoint}/{object_key}"
|
||||
return {"object_key": object_key, "url": url}
|
||||
except Exception as e:
|
||||
logger.warning(f"S3StorageClient, upload_file error: {e}")
|
||||
return {}
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
object_key: str,
|
||||
data: Union[bytes, str],
|
||||
mime: str = "application/octet-stream",
|
||||
overwrite: bool = True,
|
||||
content_disposition: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
return await make_async(self.sync_upload_file)(
|
||||
object_key, data, mime, overwrite, content_disposition
|
||||
)
|
||||
|
||||
def sync_delete_file(self, object_key: str) -> bool:
|
||||
try:
|
||||
self.client.delete_object(Bucket=self.bucket, Key=object_key)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"S3StorageClient, delete_file error: {e}")
|
||||
return False
|
||||
|
||||
async def delete_file(self, object_key: str) -> bool:
|
||||
return await make_async(self.sync_delete_file)(object_key)
|
||||
|
||||
async def close(self) -> None:
|
||||
await self.client.close()
|
||||
@@ -0,0 +1,29 @@
|
||||
import functools
|
||||
from collections import deque
|
||||
|
||||
from chainlit.context import context
|
||||
from chainlit.session import WebsocketSession
|
||||
|
||||
|
||||
def queue_until_user_message():
|
||||
def decorator(method):
|
||||
@functools.wraps(method)
|
||||
async def wrapper(self, *args, **kwargs):
|
||||
if (
|
||||
isinstance(context.session, WebsocketSession)
|
||||
and not context.session.has_first_interaction
|
||||
):
|
||||
# Queue the method invocation waiting for the first user message
|
||||
queues = context.session.thread_queues
|
||||
method_name = method.__name__
|
||||
if method_name not in queues:
|
||||
queues[method_name] = deque()
|
||||
queues[method_name].append((method, self, args, kwargs))
|
||||
|
||||
else:
|
||||
# Otherwise, Execute the method immediately
|
||||
return await method(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
Reference in New Issue
Block a user