mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-09-09 11:15:08 +08:00
Merging from main
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
]
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user