Merging from main

This commit is contained in:
Keith Stevens
2023-01-28 18:05:56 +09:00
188 changed files with 5516 additions and 2442 deletions
+29 -5
View File
@@ -1,6 +1,6 @@
from http import HTTPStatus
from secrets import token_hex
from typing import Generator
from typing import Generator, NamedTuple
from fastapi import Depends, Request, Response, Security
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
@@ -19,22 +19,46 @@ def get_db() -> Generator:
yield db
api_key_query = APIKeyQuery(name="api_key", auto_error=False)
api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)
api_key_query = APIKeyQuery(name="api_key", scheme_name="api-key", auto_error=False)
api_key_header = APIKeyHeader(name="X-API-Key", scheme_name="api-key", auto_error=False)
oasst_user_query = APIKeyQuery(name="oasst_user", scheme_name="oasst-user", auto_error=False)
oasst_user_header = APIKeyHeader(name="x-oasst-user", scheme_name="oasst-user", auto_error=False)
bearer_token = HTTPBearer(auto_error=False)
async def get_api_key(
def get_api_key(
api_key_query: str = Security(api_key_query),
api_key_header: str = Security(api_key_header),
):
) -> str:
if api_key_query:
return api_key_query
else:
return api_key_header
class FrontendUserId(NamedTuple):
auth_method: str
username: str
def get_frontend_user_id(
user_query: str = Security(oasst_user_query),
user_header: str = Security(oasst_user_header),
) -> FrontendUserId:
def split_user(v: str) -> tuple[str, str]:
if type(v) is str:
v = v.split(":", maxsplit=1)
if len(v) == 2:
return FrontendUserId(auth_method=v[0], username=v[1])
return FrontendUserId(auth_method=None, username=None)
if user_query:
return split_user(user_query)
else:
return split_user(user_header)
def create_api_client(
*,
session: Session,
+11 -5
View File
@@ -70,13 +70,14 @@ def query_frontend_user_messages(
only_roots: bool = False,
desc: bool = True,
include_deleted: bool = False,
lang: Optional[str] = None,
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Query frontend user messages.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, auth_method=auth_method, username=username)
messages = pr.query_messages_ordered_by_created_date(
auth_method=auth_method,
username=username,
@@ -87,6 +88,7 @@ def query_frontend_user_messages(
lte_created_date=end_date,
only_roots=only_roots,
deleted=None if include_deleted else False,
lang=lang,
)
return utils.prepare_message_list(messages)
@@ -95,24 +97,28 @@ def query_frontend_user_messages(
def query_frontend_user_messages_cursor(
auth_method: str,
username: str,
lt: Optional[str] = None,
gt: Optional[str] = None,
before: Optional[str] = None,
after: Optional[str] = None,
only_roots: Optional[bool] = False,
include_deleted: Optional[bool] = False,
max_count: Optional[int] = Query(10, gt=0, le=1000),
desc: Optional[bool] = False,
lang: Optional[str] = None,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
return get_messages_cursor(
lt=lt,
gt=gt,
before=before,
after=after,
auth_method=auth_method,
username=username,
only_roots=only_roots,
include_deleted=include_deleted,
max_count=max_count,
desc=desc,
lang=lang,
frontend_user=frontend_user,
api_client=api_client,
db=db,
)
+112 -32
View File
@@ -7,6 +7,7 @@ from oasst_backend.api import deps
from oasst_backend.api.v1 import utils
from oasst_backend.models import ApiClient
from oasst_backend.prompt_repository import PromptRepository
from oasst_backend.utils.database_utils import CommitMode, managed_tx_function
from oasst_shared.exceptions.oasst_api_error import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol
from sqlmodel import Session
@@ -17,6 +18,7 @@ router = APIRouter()
@router.get("/", response_model=list[protocol.Message])
def query_messages(
*,
auth_method: Optional[str] = None,
username: Optional[str] = None,
api_client_id: Optional[str] = None,
@@ -26,13 +28,15 @@ def query_messages(
only_roots: Optional[bool] = False,
desc: Optional[bool] = True,
allow_deleted: Optional[bool] = False,
lang: Optional[str] = None,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Query messages.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, auth_method=frontend_user.auth_method, username=frontend_user.username)
messages = pr.query_messages_ordered_by_created_date(
auth_method=auth_method,
username=username,
@@ -43,6 +47,7 @@ def query_messages(
lte_created_date=end_date,
only_roots=only_roots,
deleted=None if allow_deleted else False,
lang=lang,
)
return utils.prepare_message_list(messages)
@@ -50,8 +55,9 @@ def query_messages(
@router.get("/cursor", response_model=protocol.MessagePage)
def get_messages_cursor(
lt: Optional[str] = None,
gt: Optional[str] = None,
*,
before: Optional[str] = None,
after: Optional[str] = None,
user_id: Optional[UUID] = None,
auth_method: Optional[str] = None,
username: Optional[str] = None,
@@ -60,9 +66,13 @@ def get_messages_cursor(
include_deleted: Optional[bool] = False,
max_count: Optional[int] = Query(10, gt=0, le=1000),
desc: Optional[bool] = False,
lang: Optional[str] = None,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
assert max_count is not None
def split_cursor(x: str | None) -> tuple[datetime, UUID]:
if not x:
return None, None
@@ -74,11 +84,21 @@ def get_messages_cursor(
except ValueError:
raise OasstError("Invalid cursor value", OasstErrorCode.INVALID_CURSOR_VALUE)
lte_created_date, lt_id = split_cursor(lt)
gte_created_date, gt_id = split_cursor(gt)
if desc:
gte_created_date, gt_id = split_cursor(before)
lte_created_date, lt_id = split_cursor(after)
query_desc = not (before is not None and not after)
else:
lte_created_date, lt_id = split_cursor(before)
gte_created_date, gt_id = split_cursor(after)
query_desc = before is not None and not after
pr = PromptRepository(db, api_client)
messages = pr.query_messages_ordered_by_created_date(
print(f"{desc=} {query_desc=} {gte_created_date=} {lte_created_date=}")
qry_max_count = max_count + 1 if before is None or after is None else max_count
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
items = pr.query_messages_ordered_by_created_date(
user_id=user_id,
auth_method=auth_method,
username=username,
@@ -89,22 +109,31 @@ def get_messages_cursor(
lt_id=lt_id,
only_roots=only_roots,
deleted=None if include_deleted else False,
desc=desc,
limit=max_count,
desc=query_desc,
limit=qry_max_count,
lang=lang,
)
items = utils.prepare_message_list(messages)
num_rows = len(items)
if qry_max_count > max_count and num_rows == qry_max_count:
assert not (before and after)
items = items[:-1]
if desc != query_desc:
items.reverse()
items = utils.prepare_message_list(items)
n, p = None, None
if len(items) > 0:
if len(items) == max_count or gte_created_date:
if (num_rows > max_count and before) or after:
p = str(items[0].id) + "$" + items[0].created_date.isoformat()
if len(items) == max_count or lte_created_date:
if num_rows > max_count or before:
n = str(items[-1].id) + "$" + items[-1].created_date.isoformat()
else:
if gte_created_date:
p = gte_created_date.isoformat()
if lte_created_date:
n = lte_created_date.isoformat()
if after:
p = lte_created_date.isoformat() if desc else gte_created_date.isoformat()
if before:
n = gte_created_date.isoformat() if desc else lte_created_date.isoformat()
order = "desc" if desc else "asc"
return protocol.MessagePage(prev=p, next=n, sort_key="created_date", order=order, items=items)
@@ -112,37 +141,49 @@ def get_messages_cursor(
@router.get("/{message_id}", response_model=protocol.Message)
def get_message(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get a message by its internal ID.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
message = pr.fetch_message(message_id)
return utils.prepare_message(message)
@router.get("/{message_id}/conversation", response_model=protocol.Conversation)
def get_conv(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get a conversation from the tree root and up to the message with given internal ID.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
messages = pr.fetch_message_conversation(message_id)
return utils.prepare_conversation(messages)
@router.get("/{message_id}/tree", response_model=protocol.MessageTree)
def get_tree(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get all messages belonging to the same message tree.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
message = pr.fetch_message(message_id)
tree = pr.fetch_message_tree(message.message_tree_id, reviewed=False)
return utils.prepare_tree(tree, message.message_tree_id)
@@ -150,24 +191,32 @@ def get_tree(
@router.get("/{message_id}/children", response_model=list[protocol.Message])
def get_children(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get all messages belonging to the same message tree.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
messages = pr.fetch_message_children(message_id)
return utils.prepare_message_list(messages)
@router.get("/{message_id}/descendants", response_model=protocol.MessageTree)
def get_descendants(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get a subtree which starts with this message.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
message = pr.fetch_message(message_id)
descendants = pr.fetch_message_descendants(message)
return utils.prepare_tree(descendants, message.id)
@@ -175,12 +224,16 @@ def get_descendants(
@router.get("/{message_id}/longest_conversation_in_tree", response_model=protocol.Conversation)
def get_longest_conv(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get the longest conversation from the tree of the message.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
message = pr.fetch_message(message_id)
conv = pr.fetch_longest_conversation(message.message_tree_id)
return utils.prepare_conversation(conv)
@@ -188,12 +241,16 @@ def get_longest_conv(
@router.get("/{message_id}/max_children_in_tree", response_model=protocol.MessageTree)
def get_max_children(
message_id: UUID, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Get message with the most children from the tree of the provided message.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
message = pr.fetch_message(message_id)
message, children = pr.fetch_message_with_max_children(message.message_tree_id)
return utils.prepare_tree([message, *children], message.id)
@@ -201,7 +258,30 @@ def get_max_children(
@router.delete("/{message_id}", status_code=HTTP_204_NO_CONTENT)
def mark_message_deleted(
message_id: UUID, api_client: ApiClient = Depends(deps.get_trusted_api_client), db: Session = Depends(deps.get_db)
*,
message_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
db: Session = Depends(deps.get_db),
):
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
pr.mark_messages_deleted(message_id)
@router.post("/{message_id}/emoji", response_model=protocol.Message)
def post_message_emoji(
*,
message_id: UUID,
request: protocol.MessageEmojiRequest,
api_client: ApiClient = Depends(deps.get_api_client),
) -> protocol.Message:
"""
Toggle, add or remove message emoji.
"""
@managed_tx_function(CommitMode.COMMIT)
def emoji_tx(session: deps.Session):
pr = PromptRepository(session, api_client, client_user=request.user)
return pr.handle_message_emoji(message_id, request.op, request.emoji)
return utils.prepare_message(emoji_tx())
+4 -2
View File
@@ -77,6 +77,7 @@ def tasks_acknowledge(
*,
db: Session = Depends(deps.get_db),
api_key: APIKey = Depends(deps.get_api_key),
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
task_id: UUID,
ack_request: protocol_schema.TaskAck,
) -> None:
@@ -87,7 +88,7 @@ def tasks_acknowledge(
api_client = deps.api_auth(api_key, db)
try:
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
# here we store the message id in the database for the task
logger.info(f"Frontend acknowledges task {task_id=}, {ack_request=}.")
@@ -105,6 +106,7 @@ def tasks_acknowledge_failure(
*,
db: Session = Depends(deps.get_db),
api_key: APIKey = Depends(deps.get_api_key),
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
task_id: UUID,
nack_request: protocol_schema.TaskNAck,
) -> None:
@@ -115,7 +117,7 @@ def tasks_acknowledge_failure(
try:
logger.info(f"Frontend reports failure to implement task {task_id=}, {nack_request=}.")
api_client = deps.api_auth(api_key, db)
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
pr.task_repository.acknowledge_task_failure(task_id)
except (KeyError, RuntimeError):
logger.exception("Failed to not acknowledge task.")
+37 -8
View File
@@ -3,9 +3,11 @@ from fastapi.security.api_key import APIKey
from loguru import logger
from oasst_backend.api import deps
from oasst_backend.prompt_repository import PromptRepository
from oasst_backend.schemas.text_labels import LabelOption, ValidLabelsResponse
from oasst_backend.schemas.text_labels import LabelDescription, ValidLabelsResponse
from oasst_backend.utils.database_utils import CommitMode, managed_tx_function
from oasst_shared.exceptions import OasstError
from oasst_shared.schemas import protocol as protocol_schema
from sqlmodel import Session
from oasst_shared.schemas.protocol import TextLabel
from starlette.status import HTTP_204_NO_CONTENT, HTTP_400_BAD_REQUEST
router = APIRouter()
@@ -14,20 +16,25 @@ router = APIRouter()
@router.post("/", status_code=HTTP_204_NO_CONTENT)
def label_text(
*,
db: Session = Depends(deps.get_db),
api_key: APIKey = Depends(deps.get_api_key),
text_labels: protocol_schema.TextLabels,
) -> None:
"""
Label a piece of text.
"""
api_client = deps.api_auth(api_key, db)
@managed_tx_function(CommitMode.COMMIT)
def store_text_labels(session: deps.Session):
api_client = deps.api_auth(api_key, session)
pr = PromptRepository(session, api_client, client_user=text_labels.user)
pr.store_text_labels(text_labels)
try:
logger.info(f"Labeling text {text_labels=}.")
pr = PromptRepository(db, api_client, client_user=text_labels.user)
pr.store_text_labels(text_labels)
store_text_labels()
except OasstError:
raise
except Exception:
logger.exception("Failed to store label.")
raise HTTPException(
@@ -39,7 +46,29 @@ def label_text(
def get_valid_lables() -> ValidLabelsResponse:
return ValidLabelsResponse(
valid_labels=[
LabelOption(name=l.value, display_text=l.display_text, help_text=l.help_text)
for l in protocol_schema.TextLabel
LabelDescription(name=l.value, widget=l.widget.value, display_text=l.display_text, help_text=l.help_text)
for l in TextLabel
]
)
@router.get("/report_labels")
def get_report_lables() -> ValidLabelsResponse:
report_labels = [
TextLabel.spam,
TextLabel.not_appropriate,
TextLabel.pii,
TextLabel.hate_speech,
TextLabel.sexual_content,
TextLabel.moral_judgement,
TextLabel.political_content,
TextLabel.toxicity,
TextLabel.violence,
TextLabel.quality,
]
return ValidLabelsResponse(
valid_labels=[
LabelDescription(name=l.value, widget=l.widget.value, display_text=l.display_text, help_text=l.help_text)
for l in report_labels
]
)
+37 -20
View File
@@ -28,6 +28,7 @@ def get_users_ordered_by_username(
search_text: Optional[str] = None,
auth_method: Optional[str] = None,
max_count: Optional[int] = Query(100, gt=0, le=10000),
desc: Optional[bool] = False,
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
@@ -41,6 +42,7 @@ def get_users_ordered_by_username(
auth_method=auth_method,
search_text=search_text,
limit=max_count,
desc=desc,
)
return [u.to_protocol_frontend_user() for u in users]
@@ -55,6 +57,7 @@ def get_users_ordered_by_display_name(
auth_method: Optional[str] = None,
search_text: Optional[str] = None,
max_count: Optional[int] = Query(100, gt=0, le=10000),
desc: Optional[bool] = False,
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
@@ -68,14 +71,15 @@ def get_users_ordered_by_display_name(
auth_method=auth_method,
search_text=search_text,
limit=max_count,
desc=desc,
)
return [u.to_protocol_frontend_user() for u in users]
@router.get("/cursor", response_model=protocol.FrontEndUserPage)
def get_users_cursor(
lt: Optional[str] = None,
gt: Optional[str] = None,
before: Optional[str] = None,
after: Optional[str] = None,
sort_key: Optional[str] = Query("username", max_length=32),
max_count: Optional[int] = Query(100, gt=0, le=10000),
api_client_id: Optional[UUID] = None,
@@ -95,7 +99,8 @@ def get_users_cursor(
return x, None
items: list[protocol.FrontEndUser]
qry_max_count = max_count + 1 if lt is None or gt is None else max_count
qry_max_count = max_count + 1 if before is None or after is None else max_count
desc = before is not None and not after
def get_next_prev(num_rows: int, lt: str | None, gt: str | None, key_fn: Callable[[protocol.FrontEndUser], str]):
p, n = None, None
@@ -114,17 +119,16 @@ def get_users_cursor(
def remove_extra_item(items: list[protocol.FrontEndUser], lt: str | None, gt: str | None):
num_rows = len(items)
if qry_max_count > max_count and num_rows == qry_max_count:
assert not (lt and gt)
if lt:
items = items[1:]
else:
items = items[:-1]
assert not (lt is not None and gt is not None)
items = items[:-1]
if desc:
items.reverse()
return items, num_rows
n, p = None, None
if sort_key == "username":
lte_username, lt_id = split_cursor(lt)
gte_username, gt_id = split_cursor(gt)
lte_username, lt_id = split_cursor(before)
gte_username, gt_id = split_cursor(after)
items = get_users_ordered_by_username(
api_client_id=api_client_id,
gte_username=gte_username,
@@ -134,6 +138,7 @@ def get_users_cursor(
auth_method=auth_method,
search_text=search_text,
max_count=qry_max_count,
desc=desc,
api_client=api_client,
db=db,
)
@@ -141,8 +146,8 @@ def get_users_cursor(
p, n = get_next_prev(num_rows, lte_username, gte_username, lambda x: x.id)
elif sort_key == "display_name":
lte_display_name, lt_id = split_cursor(lt)
gte_display_name, gt_id = split_cursor(gt)
lte_display_name, lt_id = split_cursor(before)
gte_display_name, gt_id = split_cursor(after)
items = get_users_ordered_by_display_name(
api_client_id=api_client_id,
gte_display_name=gte_display_name,
@@ -152,6 +157,7 @@ def get_users_cursor(
auth_method=auth_method,
search_text=search_text,
max_count=qry_max_count,
desc=desc,
api_client=api_client,
db=db,
)
@@ -184,6 +190,7 @@ def update_user(
user_id: UUID,
enabled: Optional[bool] = None,
notes: Optional[str] = None,
show_on_leaderboard: Optional[bool] = None,
db: Session = Depends(deps.get_db),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
):
@@ -191,7 +198,7 @@ def update_user(
Update a user by global user ID. Only trusted clients can update users.
"""
ur = UserRepository(db, api_client)
ur.update_user(user_id, enabled, notes)
ur.update_user(user_id, enabled, notes, show_on_leaderboard)
@router.delete("/{user_id}", status_code=HTTP_204_NO_CONTENT)
@@ -217,13 +224,15 @@ def query_user_messages(
only_roots: bool = False,
desc: bool = True,
include_deleted: bool = False,
lang: Optional[str] = None,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
"""
Query user messages.
"""
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
messages = pr.query_messages_ordered_by_created_date(
user_id=user_id,
api_client_id=api_client_id,
@@ -233,6 +242,7 @@ def query_user_messages(
lte_created_date=end_date,
only_roots=only_roots,
deleted=None if include_deleted else False,
lang=lang,
)
return utils.prepare_message_list(messages)
@@ -241,23 +251,27 @@ def query_user_messages(
@router.get("/{user_id}/messages/cursor", response_model=protocol.MessagePage)
def query_user_messages_cursor(
user_id: Optional[UUID],
lt: Optional[str] = None,
gt: Optional[str] = None,
before: Optional[str] = None,
after: Optional[str] = None,
only_roots: Optional[bool] = False,
include_deleted: Optional[bool] = False,
max_count: Optional[int] = Query(10, gt=0, le=1000),
desc: Optional[bool] = False,
lang: Optional[str] = None,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_api_client),
db: Session = Depends(deps.get_db),
):
return get_messages_cursor(
lt=lt,
gt=gt,
before=before,
after=after,
user_id=user_id,
only_roots=only_roots,
include_deleted=include_deleted,
max_count=max_count,
desc=desc,
lang=lang,
frontend_user=frontend_user,
api_client=api_client,
db=db,
)
@@ -265,9 +279,12 @@ def query_user_messages_cursor(
@router.delete("/{user_id}/messages", status_code=HTTP_204_NO_CONTENT)
def mark_user_messages_deleted(
user_id: UUID, api_client: ApiClient = Depends(deps.get_trusted_api_client), db: Session = Depends(deps.get_db)
user_id: UUID,
frontend_user: deps.FrontendUserId = Depends(deps.get_frontend_user_id),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
db: Session = Depends(deps.get_db),
):
pr = PromptRepository(db, api_client)
pr = PromptRepository(db, api_client, frontend_user=frontend_user)
messages = pr.query_messages_ordered_by_created_date(user_id=user_id, limit=None)
pr.mark_messages_deleted(messages)
+15 -10
View File
@@ -14,6 +14,8 @@ def prepare_message(m: Message) -> protocol.Message:
lang=m.lang,
is_assistant=(m.role == "assistant"),
created_date=m.created_date,
emojis=m.emojis or {},
user_emojis=m.user_emojis or [],
)
@@ -21,17 +23,20 @@ def prepare_message_list(messages: list[Message]) -> list[protocol.Message]:
return [prepare_message(m) for m in messages]
def prepare_conversation_message(message: Message) -> protocol.ConversationMessage:
return protocol.ConversationMessage(
id=message.id,
frontend_message_id=message.frontend_message_id,
text=message.text,
lang=message.lang,
is_assistant=(message.role == "assistant"),
emojis=message.emojis or {},
user_emojis=message.user_emojis or [],
)
def prepare_conversation_message_list(messages: list[Message]) -> list[protocol.ConversationMessage]:
return [
protocol.ConversationMessage(
id=message.id,
frontend_message_id=message.frontend_message_id,
text=message.text,
lang=message.lang,
is_assistant=(message.role == "assistant"),
)
for message in messages
]
return [prepare_conversation_message(message) for message in messages]
def prepare_conversation(messages: list[Message]) -> protocol.Conversation: