mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-08-11 11:13:12 +08:00
Merging from main
This commit is contained in:
+34
@@ -0,0 +1,34 @@
|
||||
"""add message_id to message_reaction
|
||||
|
||||
Revision ID: 8ba17b5f467a
|
||||
Revises: 160ac010efcc
|
||||
Create Date: 2023-01-24 11:34:42.167575
|
||||
|
||||
"""
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "8ba17b5f467a"
|
||||
down_revision = "160ac010efcc"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.add_column("message_reaction", sa.Column("message_id", sqlmodel.sql.sqltypes.GUID(), nullable=True))
|
||||
op.create_index(op.f("ix_message_reaction_message_id"), "message_reaction", ["message_id"], unique=False)
|
||||
op.add_column("text_labels", sa.Column("task_id", sqlmodel.sql.sqltypes.GUID(), nullable=True))
|
||||
op.create_index(op.f("ix_text_labels_task_id"), "text_labels", ["task_id"], unique=False)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index(op.f("ix_text_labels_task_id"), table_name="text_labels")
|
||||
op.drop_column("text_labels", "task_id")
|
||||
op.drop_index(op.f("ix_message_reaction_message_id"), table_name="message_reaction")
|
||||
op.drop_column("message_reaction", "message_id")
|
||||
# ### end Alembic commands ###
|
||||
@@ -0,0 +1,44 @@
|
||||
"""add message_emoji
|
||||
|
||||
Revision ID: 40ed93df0ed5
|
||||
Revises: 8ba17b5f467a
|
||||
Create Date: 2023-01-24 22:56:28.229408
|
||||
|
||||
"""
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "40ed93df0ed5"
|
||||
down_revision = "8ba17b5f467a"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table(
|
||||
"message_emoji",
|
||||
sa.Column("message_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column(
|
||||
"created_date", sa.DateTime(timezone=True), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False
|
||||
),
|
||||
sa.Column("emoji", sqlmodel.sql.sqltypes.AutoString(length=128), nullable=False),
|
||||
sa.ForeignKeyConstraint(["message_id"], ["message.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["user.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("message_id", "user_id", "emoji"),
|
||||
)
|
||||
op.create_index("ix_message_emoji__user_id__message_id", "message_emoji", ["user_id", "message_id"], unique=False)
|
||||
op.add_column("message", sa.Column("emojis", postgresql.JSONB(astext_type=sa.Text()), nullable=True))
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_column("message", "emojis")
|
||||
op.drop_index("ix_message_emoji__user_id__message_id", table_name="message_emoji")
|
||||
op.drop_table("message_emoji")
|
||||
# ### end Alembic commands ###
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add task created date index
|
||||
|
||||
Revision ID: c84fcd6900dc
|
||||
Revises: 40ed93df0ed5
|
||||
Create Date: 2023-01-26 18:35:43.061589
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "c84fcd6900dc"
|
||||
down_revision = "40ed93df0ed5"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_index(op.f("ix_task_created_date"), "task", ["created_date"], unique=False)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index(op.f("ix_task_created_date"), table_name="task")
|
||||
# ### end Alembic commands ###
|
||||
@@ -0,0 +1,29 @@
|
||||
"""add user.show_on_leaderboard
|
||||
|
||||
Revision ID: f856bf19d32b
|
||||
Revises: c84fcd6900dc
|
||||
Create Date: 2023-01-27 20:13:56.533374
|
||||
|
||||
"""
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "f856bf19d32b"
|
||||
down_revision = "c84fcd6900dc"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.add_column(
|
||||
"user", sa.Column("show_on_leaderboard", sa.Boolean(), server_default=sa.text("true"), nullable=False)
|
||||
)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_column("user", "show_on_leaderboard")
|
||||
# ### end Alembic commands ###
|
||||
@@ -273,6 +273,38 @@ def get_openapi_schema():
|
||||
return json.dumps(app.openapi())
|
||||
|
||||
|
||||
def export_ready_trees(file: Optional[str] = None, use_compression: bool = False):
|
||||
try:
|
||||
with Session(engine) as db:
|
||||
api_client = api_auth(settings.OFFICIAL_WEB_API_KEY, db=db)
|
||||
dummy_user = protocol_schema.User(id="__dummy_user__", display_name="Dummy User", auth_method="local")
|
||||
|
||||
ur = UserRepository(db=db, api_client=api_client)
|
||||
tr = TaskRepository(db=db, api_client=api_client, client_user=dummy_user, user_repository=ur)
|
||||
pr = PromptRepository(
|
||||
db=db, api_client=api_client, client_user=dummy_user, user_repository=ur, task_repository=tr
|
||||
)
|
||||
tm = TreeManager(db, pr)
|
||||
|
||||
tm.export_all_ready_trees(file, use_compression=use_compression)
|
||||
except Exception:
|
||||
logger.exception("Error exporting trees.")
|
||||
|
||||
|
||||
def retry_scoring_failed_message_trees():
|
||||
try:
|
||||
logger.info("TreeManager.retry_scoring_failed_message_trees()")
|
||||
with Session(engine) as db:
|
||||
api_client = api_auth(settings.OFFICIAL_WEB_API_KEY, db=db)
|
||||
|
||||
pr = PromptRepository(db=db, api_client=api_client)
|
||||
tm = TreeManager(db, pr)
|
||||
tm.retry_scoring_failed_message_trees()
|
||||
|
||||
except Exception:
|
||||
logger.exception("TreeManager.retry_scoring_failed_message_trees() failed.")
|
||||
|
||||
|
||||
def main():
|
||||
# Importing here so we don't import packages unnecessarily if we're
|
||||
# importing main as a module.
|
||||
@@ -289,11 +321,28 @@ def main():
|
||||
)
|
||||
parser.add_argument("--host", help="The host to run the server", default="0.0.0.0")
|
||||
parser.add_argument("--port", help="The port to run the server", default=8080)
|
||||
parser.add_argument(
|
||||
"--export", help="Export all trees which are ready for exporting.", action=argparse.BooleanOptionalAction
|
||||
)
|
||||
parser.add_argument(
|
||||
"--export-file",
|
||||
help="Name of file to export trees to. If not provided when exporting, output will be send to STDOUT",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--retry-scoring",
|
||||
help="Retry scoring failed message trees",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.print_openapi_schema:
|
||||
print(get_openapi_schema())
|
||||
elif args.export:
|
||||
use_compression: bool = ".gz" in args.export_file
|
||||
export_ready_trees(file=args.export_file, use_compression=use_compression)
|
||||
elif args.retry_scoring:
|
||||
retry_scoring_failed_message_trees()
|
||||
else:
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from oasst_shared.schemas import protocol as protocol_schema
|
||||
from oasst_shared.schemas.protocol import TextLabel
|
||||
from pydantic import AnyHttpUrl, BaseModel, BaseSettings, FilePath, PostgresDsn, validator
|
||||
|
||||
|
||||
@@ -46,17 +46,69 @@ class TreeManagerConfiguration(BaseModel):
|
||||
num_required_rankings: int = 3
|
||||
"""Number of rankings in which the message participated."""
|
||||
|
||||
mandatory_labels_initial_prompt: Optional[list[protocol_schema.TextLabel]] = [protocol_schema.TextLabel.spam]
|
||||
labels_initial_prompt: list[TextLabel] = [
|
||||
TextLabel.spam,
|
||||
TextLabel.quality,
|
||||
TextLabel.helpfulness,
|
||||
TextLabel.creativity,
|
||||
TextLabel.humor,
|
||||
TextLabel.toxicity,
|
||||
TextLabel.violence,
|
||||
TextLabel.not_appropriate,
|
||||
TextLabel.pii,
|
||||
TextLabel.hate_speech,
|
||||
TextLabel.sexual_content,
|
||||
]
|
||||
|
||||
labels_assistant_reply: list[TextLabel] = [
|
||||
TextLabel.spam,
|
||||
TextLabel.fails_task,
|
||||
TextLabel.quality,
|
||||
TextLabel.helpfulness,
|
||||
TextLabel.creativity,
|
||||
TextLabel.humor,
|
||||
TextLabel.toxicity,
|
||||
TextLabel.violence,
|
||||
TextLabel.not_appropriate,
|
||||
TextLabel.pii,
|
||||
TextLabel.hate_speech,
|
||||
TextLabel.sexual_content,
|
||||
]
|
||||
|
||||
labels_prompter_reply: list[TextLabel] = [
|
||||
TextLabel.spam,
|
||||
TextLabel.quality,
|
||||
TextLabel.helpfulness,
|
||||
TextLabel.humor,
|
||||
TextLabel.creativity,
|
||||
TextLabel.toxicity,
|
||||
TextLabel.violence,
|
||||
TextLabel.not_appropriate,
|
||||
TextLabel.pii,
|
||||
TextLabel.hate_speech,
|
||||
TextLabel.sexual_content,
|
||||
]
|
||||
|
||||
mandatory_labels_initial_prompt: Optional[list[TextLabel]] = [TextLabel.spam]
|
||||
"""Mandatory labels in text-labeling tasks for initial prompts."""
|
||||
|
||||
mandatory_labels_assistant_reply: Optional[list[protocol_schema.TextLabel]] = [protocol_schema.TextLabel.spam]
|
||||
mandatory_labels_assistant_reply: Optional[list[TextLabel]] = [TextLabel.spam]
|
||||
"""Mandatory labels in text-labeling tasks for assistant replies."""
|
||||
|
||||
mandatory_labels_prompter_reply: Optional[list[protocol_schema.TextLabel]] = [protocol_schema.TextLabel.spam]
|
||||
mandatory_labels_prompter_reply: Optional[list[TextLabel]] = [TextLabel.spam]
|
||||
"""Mandatory labels in text-labeling tasks for prompter replies."""
|
||||
|
||||
rank_prompter_replies: bool = False
|
||||
|
||||
lonely_children_count: int = 3
|
||||
"""Number of children below which parents are preferred during sampling for reply tasks."""
|
||||
|
||||
p_lonely_child_extension: float = 0.8
|
||||
"""Probability to select a parent with less than lonely_children_count children."""
|
||||
|
||||
recent_tasks_span_sec: int = 3 * 60 # 3 min
|
||||
"""Time in seconds of recent tasks to consider for exclusion during task selection."""
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
PROJECT_NAME: str = "open-assistant backend"
|
||||
@@ -90,10 +142,13 @@ class Settings(BaseSettings):
|
||||
Path(__file__).parent.parent / "test_data/realistic/realistic_seed_data.json"
|
||||
)
|
||||
DEBUG_ALLOW_SELF_LABELING: bool = False # allow users to label their own messages
|
||||
DEBUG_ALLOW_DUPLICATE_TASKS: bool = False # offer users tasks to which they already responded
|
||||
DEBUG_SKIP_EMBEDDING_COMPUTATION: bool = False
|
||||
DEBUG_SKIP_TOXICITY_CALCULATION: bool = False
|
||||
DEBUG_DATABASE_ECHO: bool = False
|
||||
|
||||
DUPLICATE_MESSAGE_FILTER_WINDOW_MINUTES: int = 120
|
||||
|
||||
HUGGING_FACE_API_KEY: str = ""
|
||||
|
||||
ROOT_TOKENS: List[str] = ["1234"] # supply a string that can be parsed to a json list
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
__all__ = []
|
||||
@@ -1,56 +0,0 @@
|
||||
from typing import Any, Dict, Generic, List, Optional, Type, TypeVar, Union
|
||||
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from pydantic import BaseModel
|
||||
from sqlmodel import Session, SQLModel
|
||||
|
||||
ModelType = TypeVar("ModelType", bound=SQLModel)
|
||||
CreateSchemaType = TypeVar("CreateSchemaType", bound=BaseModel)
|
||||
UpdateSchemaType = TypeVar("UpdateSchemaType", bound=BaseModel)
|
||||
|
||||
|
||||
class CRUDBase(Generic[ModelType, CreateSchemaType, UpdateSchemaType]):
|
||||
def __init__(self, model: Type[ModelType]):
|
||||
"""
|
||||
CRUD object with default methods to Create, Read, Update, Delete (CRUD).
|
||||
|
||||
**Parameters**
|
||||
|
||||
* `model`: A SQLModel model class
|
||||
* `schema`: A Pydantic model (schema) class
|
||||
"""
|
||||
self.model = model
|
||||
|
||||
def get(self, db: Session, id: Any) -> Optional[ModelType]:
|
||||
return db.query(self.model).filter(self.model.id == id).first()
|
||||
|
||||
def get_multi(self, db: Session, *, begin_id: int = 0, limit: int = 100) -> List[ModelType]:
|
||||
return db.query(self.model).filter(self.model.id >= begin_id).limit(limit).all()
|
||||
|
||||
def create(self, db: Session, *, obj_in: CreateSchemaType) -> ModelType:
|
||||
obj_in_data = jsonable_encoder(obj_in)
|
||||
db_obj = self.model(**obj_in_data) # type: ignore
|
||||
db.add(db_obj)
|
||||
db.commit()
|
||||
db.refresh(db_obj)
|
||||
return db_obj
|
||||
|
||||
def update(self, db: Session, *, db_obj: ModelType, obj_in: Union[UpdateSchemaType, Dict[str, Any]]) -> ModelType:
|
||||
obj_data = jsonable_encoder(db_obj)
|
||||
if isinstance(obj_in, dict):
|
||||
update_data = obj_in
|
||||
else:
|
||||
update_data = obj_in.dict(exclude_unset=True)
|
||||
for field in obj_data:
|
||||
if field in update_data:
|
||||
setattr(db_obj, field, update_data[field])
|
||||
db.add(db_obj)
|
||||
db.commit()
|
||||
db.refresh(db_obj)
|
||||
return db_obj
|
||||
|
||||
def delete(self, db: Session, *, id: int) -> ModelType:
|
||||
obj = db.query(self.model).get(id)
|
||||
db.delete(obj)
|
||||
db.commit()
|
||||
return obj
|
||||
@@ -2,6 +2,7 @@ from .api_client import ApiClient
|
||||
from .journal import Journal, JournalIntegration
|
||||
from .message import Message
|
||||
from .message_embedding import MessageEmbedding
|
||||
from .message_emoji import MessageEmoji
|
||||
from .message_reaction import MessageReaction
|
||||
from .message_toxicity import MessageToxicity
|
||||
from .message_tree_state import MessageTreeState
|
||||
@@ -24,4 +25,5 @@ __all__ = [
|
||||
"TextLabels",
|
||||
"Journal",
|
||||
"JournalIntegration",
|
||||
"MessageEmoji",
|
||||
]
|
||||
|
||||
@@ -117,7 +117,8 @@ class LabelConversationReplyPayload(TaskPayload):
|
||||
|
||||
message_id: UUID
|
||||
conversation: protocol_schema.Conversation
|
||||
reply: str
|
||||
reply: str # deprecated
|
||||
reply_message: Optional[protocol_schema.ConversationMessage]
|
||||
valid_labels: list[str]
|
||||
mandatory_labels: Optional[list[str]]
|
||||
mode: Optional[protocol_schema.LabelTaskMode]
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
from datetime import datetime
|
||||
from http import HTTPStatus
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
from uuid import UUID, uuid4
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlalchemy.dialects.postgresql as pg
|
||||
from oasst_backend.models.db_payload import MessagePayload
|
||||
from oasst_shared.exceptions.oasst_api_error import OasstError, OasstErrorCode
|
||||
from pydantic import PrivateAttr
|
||||
from sqlalchemy import false
|
||||
from sqlmodel import Field, Index, SQLModel
|
||||
|
||||
@@ -17,6 +18,13 @@ class Message(SQLModel, table=True):
|
||||
__tablename__ = "message"
|
||||
__table_args__ = (Index("ix_message_frontend_message_id", "api_client_id", "frontend_message_id", unique=True),)
|
||||
|
||||
def __new__(cls, *args: Any, **kwargs: Any):
|
||||
new_object = super().__new__(cls, *args, **kwargs)
|
||||
# temporary fix until https://github.com/tiangolo/sqlmodel/issues/149 gets merged
|
||||
if not hasattr(new_object, "_user_emojis"):
|
||||
new_object._init_private_attributes()
|
||||
return new_object
|
||||
|
||||
id: Optional[UUID] = Field(
|
||||
sa_column=sa.Column(
|
||||
pg.UUID(as_uuid=True), primary_key=True, default=uuid4, server_default=sa.text("gen_random_uuid()")
|
||||
@@ -49,11 +57,24 @@ class Message(SQLModel, table=True):
|
||||
|
||||
rank: Optional[int] = Field(nullable=True)
|
||||
|
||||
emojis: Optional[dict[str, int]] = Field(default=None, sa_column=sa.Column(pg.JSONB), nullable=False)
|
||||
_user_emojis: Optional[list[str]] = PrivateAttr(default=None)
|
||||
|
||||
def ensure_is_message(self) -> None:
|
||||
if not self.payload or not isinstance(self.payload.payload, MessagePayload):
|
||||
raise OasstError("Invalid message", OasstErrorCode.INVALID_MESSAGE, HTTPStatus.INTERNAL_SERVER_ERROR)
|
||||
|
||||
def has_emoji(self, emoji_code: str) -> bool:
|
||||
return self.emojis and emoji_code in self.emojis and self.emojis[emoji_code] > 0
|
||||
|
||||
def has_user_emoji(self, emoji_code: str) -> bool:
|
||||
return self._user_emojis and emoji_code in self._user_emojis
|
||||
|
||||
@property
|
||||
def text(self) -> str:
|
||||
self.ensure_is_message()
|
||||
return self.payload.payload.text
|
||||
|
||||
@property
|
||||
def user_emojis(self) -> str:
|
||||
return self._user_emojis
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlalchemy.dialects.postgresql as pg
|
||||
from sqlmodel import Field, Index, SQLModel
|
||||
|
||||
|
||||
class MessageEmoji(SQLModel, table=True):
|
||||
__tablename__ = "message_emoji"
|
||||
__table_args__ = (Index("ix_message_emoji__user_id__message_id", "user_id", "message_id", unique=False),)
|
||||
|
||||
message_id: Optional[UUID] = Field(
|
||||
sa_column=sa.Column(
|
||||
pg.UUID(as_uuid=True), sa.ForeignKey("message.id", ondelete="CASCADE"), nullable=False, primary_key=True
|
||||
)
|
||||
)
|
||||
user_id: UUID = Field(
|
||||
sa_column=sa.Column(
|
||||
pg.UUID(as_uuid=True), sa.ForeignKey("user.id", ondelete="CASCADE"), nullable=False, primary_key=True
|
||||
)
|
||||
)
|
||||
emoji: str = Field(nullable=False, max_length=128, primary_key=True)
|
||||
created_date: Optional[datetime] = Field(
|
||||
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp())
|
||||
)
|
||||
@@ -26,3 +26,4 @@ class MessageReaction(SQLModel, table=True):
|
||||
payload_type: str = Field(nullable=False, max_length=200)
|
||||
payload: PayloadContainer = Field(sa_column=sa.Column(payload_column_type(PayloadContainer), nullable=False))
|
||||
api_client_id: UUID = Field(nullable=False, foreign_key="api_client.id")
|
||||
message_id: Optional[UUID] = Field(nullable=True, index=True)
|
||||
|
||||
@@ -20,7 +20,9 @@ class Task(SQLModel, table=True):
|
||||
),
|
||||
)
|
||||
created_date: Optional[datetime] = Field(
|
||||
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp()),
|
||||
sa_column=sa.Column(
|
||||
sa.DateTime(timezone=True), nullable=False, index=True, server_default=sa.func.current_timestamp()
|
||||
),
|
||||
)
|
||||
expiry_date: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(timezone=True), nullable=True))
|
||||
user_id: Optional[UUID] = Field(nullable=True, foreign_key="user.id", index=True)
|
||||
|
||||
@@ -27,3 +27,4 @@ class TextLabels(SQLModel, table=True):
|
||||
sa_column=sa.Column(pg.UUID(as_uuid=True), sa.ForeignKey("message.id"), nullable=True)
|
||||
)
|
||||
labels: dict[str, float] = Field(default={}, sa_column=sa.Column(pg.JSONB), nullable=False)
|
||||
task_id: Optional[UUID] = Field(nullable=True, index=True)
|
||||
|
||||
@@ -30,6 +30,7 @@ class User(SQLModel, table=True):
|
||||
enabled: bool = Field(sa_column=sa.Column(sa.Boolean, nullable=False, server_default=sa.true()))
|
||||
notes: str = Field(sa_column=sa.Column(AutoString(length=1024), nullable=False, server_default=""))
|
||||
deleted: bool = Field(sa_column=sa.Column(sa.Boolean, nullable=False, server_default=sa.false()))
|
||||
show_on_leaderboard: bool = Field(sa_column=sa.Column(sa.Boolean, nullable=False, server_default=sa.true()))
|
||||
|
||||
def to_protocol_frontend_user(self):
|
||||
return protocol.FrontEndUser(
|
||||
|
||||
@@ -1,18 +1,22 @@
|
||||
import random
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timedelta
|
||||
from http import HTTPStatus
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import Optional
|
||||
from uuid import UUID, uuid4
|
||||
|
||||
import oasst_backend.models.db_payload as db_payload
|
||||
import sqlalchemy as sa
|
||||
from loguru import logger
|
||||
from oasst_backend.api.deps import FrontendUserId
|
||||
from oasst_backend.config import settings
|
||||
from oasst_backend.journal_writer import JournalWriter
|
||||
from oasst_backend.models import (
|
||||
ApiClient,
|
||||
Message,
|
||||
MessageEmbedding,
|
||||
MessageEmoji,
|
||||
MessageReaction,
|
||||
MessageToxicity,
|
||||
MessageTreeState,
|
||||
@@ -28,8 +32,10 @@ from oasst_backend.utils.database_utils import CommitMode, managed_tx_method
|
||||
from oasst_shared.exceptions import OasstError, OasstErrorCode
|
||||
from oasst_shared.schemas import protocol as protocol_schema
|
||||
from oasst_shared.schemas.protocol import SystemStats
|
||||
from oasst_shared.utils import unaware_to_utc
|
||||
from sqlmodel import Session, and_, func, not_, or_, text, update
|
||||
from oasst_shared.utils import unaware_to_utc, utcnow
|
||||
from sqlalchemy.orm import Query
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
from sqlmodel import JSON, Session, and_, func, literal_column, not_, or_, text, update
|
||||
from starlette.status import HTTP_403_FORBIDDEN, HTTP_404_NOT_FOUND
|
||||
|
||||
|
||||
@@ -39,14 +45,30 @@ class PromptRepository:
|
||||
db: Session,
|
||||
api_client: ApiClient,
|
||||
client_user: Optional[protocol_schema.User] = None,
|
||||
*,
|
||||
user_repository: Optional[UserRepository] = None,
|
||||
task_repository: Optional[TaskRepository] = None,
|
||||
user_id: Optional[UUID] = None,
|
||||
auth_method: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
frontend_user: Optional[FrontendUserId] = None,
|
||||
):
|
||||
self.db = db
|
||||
self.api_client = api_client
|
||||
self.user_repository = user_repository or UserRepository(db, api_client)
|
||||
self.user = self.user_repository.lookup_client_user(client_user, create_missing=True)
|
||||
self.user_id = self.user.id if self.user else None
|
||||
|
||||
if frontend_user and not auth_method and not username:
|
||||
auth_method, username = frontend_user
|
||||
|
||||
if user_id:
|
||||
self.user = self.user_repository.get_user(id=user_id)
|
||||
self.user_id = self.user.id
|
||||
elif auth_method and username:
|
||||
self.user = self.user_repository.query_frontend_user(auth_method=auth_method, username=username)
|
||||
self.user_id = self.user.id
|
||||
else:
|
||||
self.user = self.user_repository.lookup_client_user(client_user, create_missing=True)
|
||||
self.user_id = self.user.id if self.user else None
|
||||
logger.debug(f"PromptRepository(api_client_id={self.api_client.id}, {self.user_id=})")
|
||||
self.task_repository = task_repository or TaskRepository(
|
||||
db, api_client, client_user, user_repository=self.user_repository
|
||||
@@ -168,6 +190,18 @@ class PromptRepository:
|
||||
role = None
|
||||
depth = 0
|
||||
|
||||
# reject whitespaces match with ^\s+$
|
||||
if re.match(r"^\s+$", text):
|
||||
raise OasstError("Message text is empty", OasstErrorCode.TASK_MESSAGE_TEXT_EMPTY)
|
||||
|
||||
# ensure message size is below the predefined limit
|
||||
if len(text) > settings.MESSAGE_SIZE_LIMIT:
|
||||
logger.error(f"Message size {len(text)=} exceeds size limit of {settings.MESSAGE_SIZE_LIMIT=}.")
|
||||
raise OasstError("Message size too long.", OasstErrorCode.TASK_MESSAGE_TOO_LONG)
|
||||
|
||||
if self.check_users_recent_replies_for_duplicates(text):
|
||||
raise OasstError("User recent messages have duplicates", OasstErrorCode.TASK_MESSAGE_DUPLICATED)
|
||||
|
||||
if task.parent_message_id:
|
||||
parent_message = self.fetch_message(task.parent_message_id)
|
||||
|
||||
@@ -245,7 +279,7 @@ class PromptRepository:
|
||||
|
||||
# store reaction to message
|
||||
reaction_payload = db_payload.RatingReactionPayload(rating=rating.rating)
|
||||
reaction = self.insert_reaction(message.id, reaction_payload)
|
||||
reaction = self.insert_reaction(task_id=task.id, payload=reaction_payload, message_id=message.id)
|
||||
if not task.collective:
|
||||
task.done = True
|
||||
self.db.add(task)
|
||||
@@ -255,7 +289,7 @@ class PromptRepository:
|
||||
return reaction
|
||||
|
||||
@managed_tx_method(CommitMode.COMMIT)
|
||||
def store_ranking(self, ranking: protocol_schema.MessageRanking) -> Tuple[MessageReaction, Task]:
|
||||
def store_ranking(self, ranking: protocol_schema.MessageRanking) -> tuple[MessageReaction, Task]:
|
||||
# fetch task
|
||||
task = self.task_repository.fetch_task_by_frontend_message_id(ranking.message_id)
|
||||
self._validate_task(task, frontend_message_id=ranking.message_id)
|
||||
@@ -295,7 +329,7 @@ class PromptRepository:
|
||||
ranking_parent_id=task_payload.ranking_parent_id,
|
||||
message_tree_id=task_payload.message_tree_id,
|
||||
)
|
||||
reaction = self.insert_reaction(task.id, reaction_payload)
|
||||
reaction = self.insert_reaction(task_id=task.id, payload=reaction_payload, message_id=parent_msg.id)
|
||||
self.journal.log_ranking(task, message_id=parent_msg.id, ranking=ranking.ranking)
|
||||
|
||||
logger.info(f"Ranking {ranking.ranking} stored for task {task.id}.")
|
||||
@@ -313,9 +347,8 @@ class PromptRepository:
|
||||
reaction_payload = db_payload.RankingReactionPayload(
|
||||
ranking=ranking.ranking, ranked_message_ids=ranked_message_ids
|
||||
)
|
||||
reaction = self.insert_reaction(task.id, reaction_payload)
|
||||
# TODO: resolve message_id
|
||||
self.journal.log_ranking(task, message_id=None, ranking=ranking.ranking)
|
||||
reaction = self.insert_reaction(task_id=task.id, payload=reaction_payload, message_id=None)
|
||||
# self.journal.log_ranking(task, message_id=None, ranking=ranking.ranking)
|
||||
|
||||
logger.info(f"Ranking {ranking.ranking} stored for task {task.id}.")
|
||||
|
||||
@@ -346,13 +379,13 @@ class PromptRepository:
|
||||
return message_toxicity
|
||||
|
||||
@managed_tx_method(CommitMode.FLUSH)
|
||||
def insert_message_embedding(self, message_id: UUID, model: str, embedding: List[float]) -> MessageEmbedding:
|
||||
def insert_message_embedding(self, message_id: UUID, model: str, embedding: list[float]) -> MessageEmbedding:
|
||||
"""Insert the embedding of a new message in the database.
|
||||
|
||||
Args:
|
||||
message_id (UUID): the identifier of the message we want to save its embedding
|
||||
model (str): the model used for creating the embedding
|
||||
embedding (List[float]): the values obtained from the message & model
|
||||
embedding (list[float]): the values obtained from the message & model
|
||||
|
||||
Raises:
|
||||
OasstError: if misses some of the before params
|
||||
@@ -366,7 +399,9 @@ class PromptRepository:
|
||||
return message_embedding
|
||||
|
||||
@managed_tx_method(CommitMode.FLUSH)
|
||||
def insert_reaction(self, task_id: UUID, payload: db_payload.ReactionPayload) -> MessageReaction:
|
||||
def insert_reaction(
|
||||
self, task_id: UUID, payload: db_payload.ReactionPayload, message_id: Optional[UUID]
|
||||
) -> MessageReaction:
|
||||
self.ensure_user_is_enabled()
|
||||
|
||||
container = PayloadContainer(payload=payload)
|
||||
@@ -376,12 +411,13 @@ class PromptRepository:
|
||||
payload=container,
|
||||
api_client_id=self.api_client.id,
|
||||
payload_type=type(payload).__name__,
|
||||
message_id=message_id,
|
||||
)
|
||||
self.db.add(reaction)
|
||||
return reaction
|
||||
|
||||
@managed_tx_method(CommitMode.FLUSH)
|
||||
def store_text_labels(self, text_labels: protocol_schema.TextLabels) -> Tuple[TextLabels, Task, Message]:
|
||||
def store_text_labels(self, text_labels: protocol_schema.TextLabels) -> tuple[TextLabels, Task, Message]:
|
||||
|
||||
valid_labels: Optional[list[str]] = None
|
||||
mandatory_labels: Optional[list[str]] = None
|
||||
@@ -441,11 +477,24 @@ class PromptRepository:
|
||||
user_id=self.user_id,
|
||||
text=text_labels.text,
|
||||
labels=text_labels.labels,
|
||||
task_id=task.id if task else None,
|
||||
)
|
||||
|
||||
if message_id:
|
||||
message = self.fetch_message(message_id)
|
||||
if task:
|
||||
if not task:
|
||||
if text_labels.is_report is True:
|
||||
message = self.handle_message_emoji(
|
||||
message_id, protocol_schema.EmojiOp.add, protocol_schema.EmojiCode.red_flag
|
||||
)
|
||||
|
||||
# update existing record for repeated updates (same user no task associated)
|
||||
existing_text_label = self.fetch_non_task_text_labels(message_id, self.user_id)
|
||||
if existing_text_label is not None:
|
||||
existing_text_label.labels = text_labels.labels
|
||||
model = existing_text_label
|
||||
|
||||
else:
|
||||
message = self.fetch_message(message_id)
|
||||
message.review_count += 1
|
||||
self.db.add(message)
|
||||
|
||||
@@ -519,6 +568,46 @@ class PromptRepository:
|
||||
qry = qry.filter(Message.review_result)
|
||||
if not include_deleted:
|
||||
qry = qry.filter(not_(Message.deleted))
|
||||
return self._add_user_emojis_all(qry)
|
||||
|
||||
def check_users_recent_replies_for_duplicates(self, text: str) -> bool:
|
||||
"""
|
||||
Checks if the user has recently replied with the same text within a given time period.
|
||||
"""
|
||||
|
||||
user_id = self.user_id
|
||||
logger.debug(f"Checking for duplicate tasks for user {user_id}")
|
||||
# messages in the past 24 hours
|
||||
messages = (
|
||||
self.db.query(Message)
|
||||
.filter(Message.user_id == user_id)
|
||||
.order_by(Message.created_date.desc())
|
||||
.filter(
|
||||
Message.created_date > utcnow() - timedelta(minutes=settings.DUPLICATE_MESSAGE_FILTER_WINDOW_MINUTES)
|
||||
)
|
||||
.all()
|
||||
)
|
||||
if not messages:
|
||||
return False
|
||||
for msg in messages:
|
||||
if msg.text == text:
|
||||
return True
|
||||
return False
|
||||
|
||||
def fetch_user_message_trees(
|
||||
self, user_id: Message.user_id, reviewed: bool = True, include_deleted: bool = False
|
||||
) -> list[Message]:
|
||||
qry = self.db.query(Message).filter(Message.user_id == user_id)
|
||||
if reviewed:
|
||||
qry = qry.filter(Message.review_result)
|
||||
if not include_deleted:
|
||||
qry = qry.filter(not_(Message.deleted))
|
||||
return self._add_user_emojis_all(qry)
|
||||
|
||||
def fetch_message_trees_ready_for_export(self) -> list[MessageTreeState]:
|
||||
qry = self.db.query(MessageTreeState).filter(
|
||||
MessageTreeState.state == message_tree_state.State.READY_FOR_EXPORT
|
||||
)
|
||||
return qry.all()
|
||||
|
||||
def fetch_multiple_random_replies(self, max_size: int = 5, message_role: str = None):
|
||||
@@ -556,11 +645,25 @@ class PromptRepository:
|
||||
return conversation, replies
|
||||
|
||||
def fetch_message(self, message_id: UUID, fail_if_missing: bool = True) -> Optional[Message]:
|
||||
qry = self.db.query(Message).filter(Message.id == message_id)
|
||||
messages = self._add_user_emojis_all(qry)
|
||||
message = messages[0] if messages else None
|
||||
|
||||
message = self.db.query(Message).filter(Message.id == message_id).one_or_none()
|
||||
if fail_if_missing and not message:
|
||||
raise OasstError("Message not found", OasstErrorCode.MESSAGE_NOT_FOUND, HTTP_404_NOT_FOUND)
|
||||
return message
|
||||
|
||||
def fetch_non_task_text_labels(self, message_id: UUID, user_id: UUID) -> Optional[TextLabels]:
|
||||
|
||||
query = (
|
||||
self.db.query(TextLabels)
|
||||
.outerjoin(Task, Task.id == TextLabels.id)
|
||||
.filter(Task.id.is_(None), TextLabels.message_id == message_id, TextLabels.user_id == user_id)
|
||||
)
|
||||
text_label = query.one_or_none()
|
||||
return text_label
|
||||
|
||||
@staticmethod
|
||||
def trace_conversation(messages: list[Message] | dict[UUID, Message], last_message: Message) -> list[Message]:
|
||||
"""
|
||||
@@ -620,9 +723,27 @@ class PromptRepository:
|
||||
qry = qry.filter(Message.review_result)
|
||||
if exclude_deleted:
|
||||
qry = qry.filter(Message.deleted == sa.false())
|
||||
children = qry.all()
|
||||
children = self._add_user_emojis_all(qry)
|
||||
return children
|
||||
|
||||
def fetch_message_siblings(
|
||||
self, message: Message | UUID, reviewed: Optional[bool] = True, deleted: Optional[bool] = False
|
||||
) -> list[Message]:
|
||||
"""
|
||||
Get siblings of a message (other messages with the same parent_id)
|
||||
"""
|
||||
if isinstance(message, Message):
|
||||
message = message.id
|
||||
|
||||
parent_qry = self.db.query(Message.parent_id).filter(Message.id == message).subquery()
|
||||
qry = self.db.query(Message).filter(Message.parent_id == parent_qry.c.parent_id)
|
||||
if reviewed is not None:
|
||||
qry = qry.filter(Message.review_result == reviewed)
|
||||
if deleted is not None:
|
||||
qry = qry.filter(Message.deleted == deleted)
|
||||
siblings = self._add_user_emojis_all(qry)
|
||||
return siblings
|
||||
|
||||
@staticmethod
|
||||
def trace_descendants(root: Message, messages: list[Message]) -> list[Message]:
|
||||
children = defaultdict(list)
|
||||
@@ -651,7 +772,7 @@ class PromptRepository:
|
||||
if max_depth is not None:
|
||||
desc = desc.filter(Message.depth <= max_depth)
|
||||
|
||||
desc = desc.all()
|
||||
desc = self._add_user_emojis_all(desc)
|
||||
|
||||
return self.trace_descendants(message, desc)
|
||||
|
||||
@@ -665,6 +786,33 @@ class PromptRepository:
|
||||
max_message = max(tree, key=lambda m: m.children_count)
|
||||
return max_message, [m for m in tree if m.parent_id == max_message.id]
|
||||
|
||||
def _add_user_emojis_all(self, qry: Query) -> list[Message]:
|
||||
if self.user_id is None:
|
||||
return qry.all()
|
||||
|
||||
sq = qry.subquery("m")
|
||||
qry = (
|
||||
self.db.query(Message, func.string_agg(MessageEmoji.emoji, literal_column("','")).label("user_emojis"))
|
||||
.select_entity_from(sq)
|
||||
.outerjoin(
|
||||
MessageEmoji,
|
||||
and_(
|
||||
sq.c.id == MessageEmoji.message_id,
|
||||
MessageEmoji.user_id == self.user_id,
|
||||
sq.c.emojis != JSON.NULL,
|
||||
),
|
||||
)
|
||||
.group_by(sq)
|
||||
)
|
||||
messages: list[Message] = []
|
||||
for x in qry:
|
||||
m: Message = x.Message
|
||||
user_emojis = x["user_emojis"]
|
||||
if user_emojis:
|
||||
m._user_emojis = user_emojis.split(",")
|
||||
messages.append(m)
|
||||
return messages
|
||||
|
||||
def query_messages_ordered_by_created_date(
|
||||
self,
|
||||
user_id: Optional[UUID] = None,
|
||||
@@ -679,6 +827,7 @@ class PromptRepository:
|
||||
deleted: Optional[bool] = None,
|
||||
desc: bool = False,
|
||||
limit: Optional[int] = 100,
|
||||
lang: Optional[str] = None,
|
||||
) -> list[Message]:
|
||||
if not self.api_client.trusted:
|
||||
if not api_client_id:
|
||||
@@ -693,7 +842,7 @@ class PromptRepository:
|
||||
if user_id:
|
||||
qry = qry.filter(Message.user_id == user_id)
|
||||
if username or auth_method:
|
||||
if not username and auth_method:
|
||||
if not (username and auth_method):
|
||||
raise OasstError("Auth method or username missing.", OasstErrorCode.AUTH_AND_USERNAME_REQUIRED)
|
||||
qry = qry.join(User)
|
||||
qry = qry.filter(User.username == username, User.auth_method == auth_method)
|
||||
@@ -743,7 +892,10 @@ class PromptRepository:
|
||||
if limit is not None:
|
||||
qry = qry.limit(limit)
|
||||
|
||||
return qry.all()
|
||||
if lang is not None:
|
||||
qry = qry.filter(Message.lang == lang)
|
||||
|
||||
return self._add_user_emojis_all(qry)
|
||||
|
||||
def update_children_counts(self, message_tree_id: UUID):
|
||||
sql_update_children_count = """
|
||||
@@ -805,3 +957,82 @@ WHERE message.id = cc.id;
|
||||
deleted=result.get(True, 0),
|
||||
message_trees=result.get(None, 0),
|
||||
)
|
||||
|
||||
def handle_message_emoji(self, message_id: UUID, op: protocol_schema.EmojiOp, emoji: protocol_schema) -> Message:
|
||||
self.ensure_user_is_enabled()
|
||||
|
||||
message = self.fetch_message(message_id)
|
||||
|
||||
# check if emoji exists
|
||||
existing_emoji = (
|
||||
self.db.query(MessageEmoji)
|
||||
.filter(
|
||||
MessageEmoji.message_id == message_id, MessageEmoji.user_id == self.user_id, MessageEmoji.emoji == emoji
|
||||
)
|
||||
.one_or_none()
|
||||
)
|
||||
|
||||
if existing_emoji:
|
||||
if op == protocol_schema.EmojiOp.add:
|
||||
logger.info(f"Emoji record already exists {message_id=}, {emoji=}, {self.user_id=}")
|
||||
return message
|
||||
elif op == protocol_schema.EmojiOp.togggle:
|
||||
op = protocol_schema.EmojiOp.remove
|
||||
|
||||
if existing_emoji is None:
|
||||
if op == protocol_schema.EmojiOp.remove:
|
||||
logger.info(f"Emoji record not found {message_id=}, {emoji=}, {self.user_id=}")
|
||||
return message
|
||||
elif op == protocol_schema.EmojiOp.togggle:
|
||||
op = protocol_schema.EmojiOp.add
|
||||
|
||||
if op == protocol_schema.EmojiOp.add:
|
||||
# hard coded exclusivity of thumbs_up & thumbs_down
|
||||
if emoji == protocol_schema.EmojiCode.thumbs_up and message.has_user_emoji(
|
||||
protocol_schema.EmojiCode.thumbs_down.value
|
||||
):
|
||||
message = self.handle_message_emoji(
|
||||
message_id, protocol_schema.EmojiOp.remove, protocol_schema.EmojiCode.thumbs_down
|
||||
)
|
||||
elif emoji == protocol_schema.EmojiCode.thumbs_down and message.has_user_emoji(
|
||||
protocol_schema.EmojiCode.thumbs_up.value
|
||||
):
|
||||
message = self.handle_message_emoji(
|
||||
message_id, protocol_schema.EmojiOp.remove, protocol_schema.EmojiCode.thumbs_up
|
||||
)
|
||||
|
||||
# insert emoji record & increment count
|
||||
message_emoji = MessageEmoji(message_id=message.id, user_id=self.user_id, emoji=emoji)
|
||||
self.db.add(message_emoji)
|
||||
emoji_counts = message.emojis
|
||||
if not emoji_counts:
|
||||
message.emojis = {emoji.value: 1}
|
||||
else:
|
||||
count = emoji_counts.get(emoji.value) or 0
|
||||
emoji_counts[emoji.value] = count + 1
|
||||
if message._user_emojis is None:
|
||||
message._user_emojis = []
|
||||
if emoji.value not in message._user_emojis:
|
||||
message._user_emojis.append(emoji.value)
|
||||
elif op == protocol_schema.EmojiOp.remove:
|
||||
# remove emoji record and & decrement count
|
||||
message = self.fetch_message(message_id)
|
||||
if message._user_emojis and emoji.value in message._user_emojis:
|
||||
message._user_emojis.remove(emoji.value)
|
||||
self.db.delete(existing_emoji)
|
||||
emoji_counts = message.emojis
|
||||
count = emoji_counts.get(emoji.value)
|
||||
if count is not None:
|
||||
if count == 1:
|
||||
del emoji_counts[emoji.value]
|
||||
else:
|
||||
emoji_counts[emoji.value] = count - 1
|
||||
flag_modified(message, "emojis")
|
||||
self.db.add(message)
|
||||
else:
|
||||
raise OasstError("Emoji op not supported", OasstErrorCode.EMOJI_OP_UNSUPPORTED)
|
||||
|
||||
flag_modified(message, "emojis")
|
||||
self.db.add(message)
|
||||
self.db.flush()
|
||||
return message
|
||||
|
||||
@@ -1,13 +1,6 @@
|
||||
from typing import Optional
|
||||
|
||||
from oasst_shared.schemas.protocol import LabelDescription
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class LabelOption(BaseModel):
|
||||
name: str
|
||||
display_text: str
|
||||
help_text: Optional[str]
|
||||
|
||||
|
||||
class ValidLabelsResponse(BaseModel):
|
||||
valid_labels: list[LabelOption]
|
||||
valid_labels: list[LabelDescription]
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from datetime import timedelta
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
@@ -9,7 +10,7 @@ from oasst_backend.user_repository import UserRepository
|
||||
from oasst_backend.utils.database_utils import CommitMode, managed_tx_method
|
||||
from oasst_shared.exceptions.oasst_api_error import OasstError, OasstErrorCode
|
||||
from oasst_shared.schemas import protocol as protocol_schema
|
||||
from sqlmodel import Session
|
||||
from sqlmodel import Session, func, or_
|
||||
from starlette.status import HTTP_404_NOT_FOUND
|
||||
|
||||
|
||||
@@ -100,6 +101,7 @@ class TaskRepository:
|
||||
message_id=task.message_id,
|
||||
conversation=task.conversation,
|
||||
reply=task.reply,
|
||||
reply_message=task.reply_message,
|
||||
valid_labels=task.valid_labels,
|
||||
mandatory_labels=task.mandatory_labels,
|
||||
mode=task.mode,
|
||||
@@ -111,6 +113,7 @@ class TaskRepository:
|
||||
message_id=task.message_id,
|
||||
conversation=task.conversation,
|
||||
reply=task.reply,
|
||||
reply_message=task.reply_message,
|
||||
valid_labels=task.valid_labels,
|
||||
mandatory_labels=task.mandatory_labels,
|
||||
mode=task.mode,
|
||||
@@ -219,3 +222,16 @@ class TaskRepository:
|
||||
def fetch_task_by_id(self, task_id: UUID) -> Task:
|
||||
task = self.db.query(Task).filter(Task.api_client_id == self.api_client.id, Task.id == task_id).one_or_none()
|
||||
return task
|
||||
|
||||
def fetch_recent_reply_tasks(
|
||||
self, max_age: timedelta = timedelta(minutes=5), done: bool = False, limit: int = 100
|
||||
) -> list[Task]:
|
||||
qry = self.db.query(Task).filter(
|
||||
func.age(Task.created_date) < max_age,
|
||||
or_(Task.payload_type == "AssistantReplyPayload", Task.payload_type == "PrompterReplyPayload"),
|
||||
)
|
||||
if done is not None:
|
||||
qry = qry.filter(Task.done == done)
|
||||
if limit:
|
||||
qry = qry.limit(limit)
|
||||
return qry.all()
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import json
|
||||
import random
|
||||
from datetime import datetime
|
||||
import sys
|
||||
from datetime import datetime, timedelta
|
||||
from enum import Enum
|
||||
from http import HTTPStatus
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
@@ -7,11 +9,17 @@ from uuid import UUID
|
||||
|
||||
import numpy as np
|
||||
import pydantic
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from loguru import logger
|
||||
from oasst_backend.api.v1.utils import prepare_conversation, prepare_conversation_message_list
|
||||
from oasst_backend.api.v1.utils import (
|
||||
prepare_conversation,
|
||||
prepare_conversation_message,
|
||||
prepare_conversation_message_list,
|
||||
)
|
||||
from oasst_backend.config import TreeManagerConfiguration, settings
|
||||
from oasst_backend.models import Message, MessageReaction, MessageTreeState, Task, TextLabels, User, message_tree_state
|
||||
from oasst_backend.prompt_repository import PromptRepository
|
||||
from oasst_backend.utils import tree_export
|
||||
from oasst_backend.utils.database_utils import CommitMode, async_managed_tx_method, managed_tx_method
|
||||
from oasst_backend.utils.hugging_face import HfClassificationModel, HfEmbeddingModel, HfUrl, HuggingFaceAPI
|
||||
from oasst_backend.utils.ranking import ranked_pairs
|
||||
@@ -37,8 +45,9 @@ class TaskRole(Enum):
|
||||
|
||||
class ActiveTreeSizeRow(pydantic.BaseModel):
|
||||
message_tree_id: UUID
|
||||
tree_size: int
|
||||
goal_tree_size: int
|
||||
tree_size: int
|
||||
awaiting_review: Optional[int]
|
||||
|
||||
@property
|
||||
def remaining_messages(self) -> int:
|
||||
@@ -89,8 +98,6 @@ class TreeManagerStats(pydantic.BaseModel):
|
||||
|
||||
|
||||
class TreeManager:
|
||||
_all_text_labels = list(map(lambda x: x.value, protocol_schema.TextLabel))
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
@@ -197,7 +204,7 @@ class TreeManager:
|
||||
lang = "en"
|
||||
logger.warning("Task availability request without lang tag received, assuming lang='en'.")
|
||||
|
||||
num_active_trees = self.query_num_active_trees(lang=lang)
|
||||
num_active_trees = self.query_num_active_trees(lang=lang, exclude_ranking=True)
|
||||
extendible_parents = self.query_extendible_parents(lang=lang)
|
||||
prompts_need_review = self.query_prompts_need_review(lang=lang)
|
||||
replies_need_review = self.query_replies_need_review(lang=lang)
|
||||
@@ -211,6 +218,15 @@ class TreeManager:
|
||||
incomplete_rankings=incomplete_rankings,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_label_descriptions(valid_labels: list[TextLabels]) -> list[protocol_schema.LabelDescription]:
|
||||
return [
|
||||
protocol_schema.LabelDescription(
|
||||
name=l.value, widget=l.widget.value, display_text=l.display_text, help_text=l.help_text
|
||||
)
|
||||
for l in valid_labels
|
||||
]
|
||||
|
||||
def next_task(
|
||||
self,
|
||||
desired_task_type: protocol_schema.TaskRequestType = protocol_schema.TaskRequestType.random,
|
||||
@@ -225,7 +241,7 @@ class TreeManager:
|
||||
lang = "en"
|
||||
logger.warning("Task request without lang tag received, assuming 'en'.")
|
||||
|
||||
num_active_trees = self.query_num_active_trees(lang=lang)
|
||||
num_active_trees = self.query_num_active_trees(lang=lang, exclude_ranking=True)
|
||||
prompts_need_review = self.query_prompts_need_review(lang=lang)
|
||||
replies_need_review = self.query_replies_need_review(lang=lang)
|
||||
extendible_parents = self.query_extendible_parents(lang=lang)
|
||||
@@ -334,6 +350,7 @@ class TreeManager:
|
||||
message_tree_id = messages[-1].message_tree_id
|
||||
|
||||
case TaskType.LABEL_REPLY:
|
||||
|
||||
if task_role == TaskRole.PROMPTER:
|
||||
replies_need_review = list(filter(lambda m: m.role == "prompter", replies_need_review))
|
||||
elif task_role == TaskRole.ASSISTANT:
|
||||
@@ -349,58 +366,99 @@ class TreeManager:
|
||||
self.cfg.p_full_labeling_review_reply_prompter: float = 0.1
|
||||
|
||||
label_mode = protocol_schema.LabelTaskMode.full
|
||||
valid_labels = self._all_text_labels
|
||||
label_disposition = protocol_schema.LabelTaskDisposition.quality
|
||||
|
||||
if message.role == "assistant":
|
||||
valid_labels = self.cfg.labels_assistant_reply
|
||||
if (
|
||||
desired_task_type == protocol_schema.TaskRequestType.random
|
||||
and random.random() > self.cfg.p_full_labeling_review_reply_assistant
|
||||
):
|
||||
valid_labels = list(map(lambda x: x.value, self.cfg.mandatory_labels_assistant_reply))
|
||||
label_mode = protocol_schema.LabelTaskMode.simple
|
||||
label_disposition = protocol_schema.LabelTaskDisposition.spam
|
||||
valid_labels = list(self.cfg.mandatory_labels_assistant_reply)
|
||||
if protocol_schema.TextLabel.quality not in valid_labels:
|
||||
valid_labels.append(protocol_schema.TextLabel.quality)
|
||||
|
||||
logger.info(f"Generating a LabelAssistantReplyTask. ({label_mode=:s})")
|
||||
task = protocol_schema.LabelAssistantReplyTask(
|
||||
message_id=message.id,
|
||||
conversation=conversation,
|
||||
reply=message.text,
|
||||
valid_labels=valid_labels,
|
||||
reply_message=prepare_conversation_message(message),
|
||||
valid_labels=list(map(lambda x: x.value, valid_labels)),
|
||||
mandatory_labels=list(map(lambda x: x.value, self.cfg.mandatory_labels_assistant_reply)),
|
||||
mode=label_mode,
|
||||
disposition=label_disposition,
|
||||
labels=self._get_label_descriptions(valid_labels),
|
||||
)
|
||||
else:
|
||||
valid_labels = self.cfg.labels_prompter_reply
|
||||
if (
|
||||
desired_task_type == protocol_schema.TaskRequestType.random
|
||||
and random.random() > self.cfg.p_full_labeling_review_reply_prompter
|
||||
):
|
||||
valid_labels = list(map(lambda x: x.value, self.cfg.mandatory_labels_prompter_reply))
|
||||
label_mode = protocol_schema.LabelTaskMode.simple
|
||||
label_disposition = protocol_schema.LabelTaskDisposition.spam
|
||||
valid_labels = list(self.cfg.mandatory_labels_prompter_reply)
|
||||
if protocol_schema.TextLabel.quality not in valid_labels:
|
||||
valid_labels.append(protocol_schema.TextLabel.quality)
|
||||
|
||||
logger.info(f"Generating a LabelPrompterReplyTask. ({label_mode=:s})")
|
||||
task = protocol_schema.LabelPrompterReplyTask(
|
||||
message_id=message.id,
|
||||
conversation=conversation,
|
||||
reply=message.text,
|
||||
valid_labels=valid_labels,
|
||||
reply_message=prepare_conversation_message(message),
|
||||
valid_labels=list(map(lambda x: x.value, valid_labels)),
|
||||
mandatory_labels=list(map(lambda x: x.value, self.cfg.mandatory_labels_prompter_reply)),
|
||||
mode=label_mode,
|
||||
disposition=label_disposition,
|
||||
labels=self._get_label_descriptions(valid_labels),
|
||||
)
|
||||
|
||||
parent_message_id = message.id
|
||||
message_tree_id = message.message_tree_id
|
||||
|
||||
case TaskType.REPLY:
|
||||
# select a tree with missing replies
|
||||
|
||||
recent_reply_tasks = self.pr.task_repository.fetch_recent_reply_tasks(
|
||||
max_age=timedelta(seconds=self.cfg.recent_tasks_span_sec), done=False
|
||||
)
|
||||
recent_reply_task_parents = {t.parent_message_id for t in recent_reply_tasks}
|
||||
|
||||
if task_role == TaskRole.PROMPTER:
|
||||
extendible_parents = list(filter(lambda x: x.parent_role == "assistant", extendible_parents))
|
||||
elif task_role == TaskRole.ASSISTANT:
|
||||
extendible_parents = list(filter(lambda x: x.parent_role == "prompter", extendible_parents))
|
||||
|
||||
# select a tree with missing replies
|
||||
if len(extendible_parents) > 0:
|
||||
random_parent = random.choice(extendible_parents)
|
||||
random_parent: ExtendibleParentRow = None
|
||||
if self.cfg.p_lonely_child_extension > 0 and self.cfg.lonely_children_count > 1:
|
||||
# check if we have extendible parents with a small number of replies
|
||||
|
||||
lonely_children_parents = [
|
||||
p
|
||||
for p in extendible_parents
|
||||
if 0 < p.active_children_count < self.cfg.lonely_children_count
|
||||
and p.parent_id not in recent_reply_task_parents
|
||||
]
|
||||
if len(lonely_children_parents) > 0 and random.random() < self.cfg.p_lonely_child_extension:
|
||||
random_parent = random.choice(lonely_children_parents)
|
||||
|
||||
if random_parent is None:
|
||||
# try to exclude parents for which tasks were recently handed out
|
||||
fresh_parents = [p for p in extendible_parents if p.parent_id not in recent_reply_task_parents]
|
||||
if len(fresh_parents) > 0:
|
||||
random_parent = random.choice(fresh_parents)
|
||||
else:
|
||||
random_parent = random.choice(extendible_parents)
|
||||
|
||||
# fetch random conversation to extend
|
||||
logger.debug(f"selected {random_parent=}")
|
||||
messages = self.pr.fetch_message_conversation(random_parent.parent_id)
|
||||
assert all(m.review_result for m in messages) # ensure all messages have positive review
|
||||
assert all(m.review_result for m in messages) # ensure all messages have positive reviews
|
||||
conversation = prepare_conversation(messages)
|
||||
|
||||
# generate reply task depending on last message
|
||||
@@ -419,19 +477,23 @@ class TreeManager:
|
||||
message = random.choice(prompts_need_review)
|
||||
|
||||
label_mode = protocol_schema.LabelTaskMode.full
|
||||
valid_labels = self._all_text_labels
|
||||
label_disposition = protocol_schema.LabelTaskDisposition.quality
|
||||
valid_labels = self.cfg.labels_initial_prompt
|
||||
|
||||
if random.random() > self.cfg.p_full_labeling_review_prompt:
|
||||
valid_labels = list(map(lambda x: x.value, self.cfg.mandatory_labels_initial_prompt))
|
||||
valid_labels = self.cfg.mandatory_labels_initial_prompt
|
||||
label_mode = protocol_schema.LabelTaskMode.simple
|
||||
label_disposition = protocol_schema.LabelTaskDisposition.spam
|
||||
|
||||
logger.info(f"Generating a LabelInitialPromptTask ({label_mode=:s}).")
|
||||
task = protocol_schema.LabelInitialPromptTask(
|
||||
message_id=message.id,
|
||||
prompt=message.text,
|
||||
valid_labels=valid_labels,
|
||||
valid_labels=list(map(lambda x: x.value, valid_labels)),
|
||||
mandatory_labels=list(map(lambda x: x.value, self.cfg.mandatory_labels_initial_prompt)),
|
||||
mode=label_mode,
|
||||
disposition=label_disposition,
|
||||
labels=self._get_label_descriptions(valid_labels),
|
||||
)
|
||||
|
||||
parent_message_id = message.id
|
||||
@@ -464,14 +526,6 @@ class TreeManager:
|
||||
logger.info(
|
||||
f"Frontend reports text reply to {interaction.message_id=} with {interaction.text=} by {interaction.user=}."
|
||||
)
|
||||
|
||||
# ensure message size is below the predefined limit
|
||||
if len(interaction.text) > settings.MESSAGE_SIZE_LIMIT:
|
||||
logger.error(
|
||||
f"Message size {len(interaction.text)=} exceeds size limit of {settings.MESSAGE_SIZE_LIMIT=}."
|
||||
)
|
||||
raise OasstError("Message size too long.", OasstErrorCode.TASK_MESSAGE_TOO_LONG)
|
||||
|
||||
# here we store the text reply in the database
|
||||
message = pr.store_text_reply(
|
||||
text=interaction.text,
|
||||
@@ -502,19 +556,18 @@ class TreeManager:
|
||||
try:
|
||||
model_name: str = HfClassificationModel.TOXIC_ROBERTA.value
|
||||
hugging_face_api: HuggingFaceAPI = HuggingFaceAPI(
|
||||
f"{HfUrl.HUGGINGFACE_FEATURE_EXTRACTION.value}/{model_name}"
|
||||
f"{HfUrl.HUGGINGFACE_TOXIC_CLASSIFICATION.value}/{model_name}"
|
||||
)
|
||||
|
||||
toxicity: List[List[Dict[str, Any]]] = await hugging_face_api.post(interaction.text)
|
||||
toxicity = toxicity[0][0]
|
||||
|
||||
pr.insert_toxicity(
|
||||
message_id=message.id, model=model_name, score=toxicity["score"], label=toxicity["label"]
|
||||
)
|
||||
|
||||
except OasstError:
|
||||
logger.error(
|
||||
f"Could not compute toxicity for text reply to {interaction.message_id=} with {interaction.text=} by {interaction.user=}."
|
||||
f"Could not compute toxicity for text reply to {interaction.message_id=} with {interaction.text=} by {interaction.user=}."
|
||||
)
|
||||
|
||||
case protocol_schema.MessageRating:
|
||||
@@ -530,9 +583,7 @@ class TreeManager:
|
||||
)
|
||||
|
||||
_, task = pr.store_ranking(interaction)
|
||||
|
||||
ok, rankings_by_message = self.check_condition_for_scoring_state(task.message_tree_id)
|
||||
self.update_message_ranks(task.message_tree_id, rankings_by_message)
|
||||
self.check_condition_for_scoring_state(task.message_tree_id)
|
||||
|
||||
case protocol_schema.TextLabels:
|
||||
logger.info(
|
||||
@@ -541,7 +592,7 @@ class TreeManager:
|
||||
|
||||
_, task, msg = pr.store_text_labels(interaction)
|
||||
|
||||
# if it was a respones for a task, check if we have enough reviews to calc review_result
|
||||
# if it was a response for a task, check if we have enough reviews to calc review_result
|
||||
if task and msg:
|
||||
reviews = self.query_reviews_for_message(msg.id)
|
||||
acceptance_score = self._calculate_acceptance(reviews)
|
||||
@@ -622,8 +673,8 @@ class TreeManager:
|
||||
|
||||
# check if desired tree size has been reached and all nodes have been reviewed
|
||||
tree_size = self.query_tree_size(message_tree_id)
|
||||
if tree_size.remaining_messages > 0:
|
||||
logger.debug(f"False {tree_size.remaining_messages=}")
|
||||
if tree_size.remaining_messages > 0 or tree_size.awaiting_review > 0:
|
||||
logger.debug(f"False {tree_size.remaining_messages=}, {tree_size.awaiting_review=}")
|
||||
return False
|
||||
|
||||
self._enter_state(mts, message_tree_state.State.RANKING)
|
||||
@@ -647,9 +698,12 @@ class TreeManager:
|
||||
return False, None
|
||||
|
||||
self._enter_state(mts, message_tree_state.State.READY_FOR_SCORING)
|
||||
return True, rankings_by_message
|
||||
self.update_message_ranks(message_tree_id, rankings_by_message)
|
||||
return True
|
||||
|
||||
def update_message_ranks(self, message_tree_id: UUID, rankings_by_message: Dict[int, int]) -> bool:
|
||||
def update_message_ranks(
|
||||
self, message_tree_id: UUID, rankings_by_message: dict[UUID, list[MessageReaction]]
|
||||
) -> bool:
|
||||
|
||||
mts = self.pr.fetch_tree_state(message_tree_id)
|
||||
# check state, allow retry if in SCORING_FAILED state
|
||||
@@ -657,19 +711,47 @@ class TreeManager:
|
||||
logger.debug(f"False {mts.active=}, {mts.state=}")
|
||||
return False
|
||||
|
||||
if mts.state == message_tree_state.State.SCORING_FAILED:
|
||||
mts.active = True
|
||||
mts.state = message_tree_state.State.READY_FOR_SCORING
|
||||
|
||||
try:
|
||||
for rankings in rankings_by_message.values():
|
||||
sorted_messages = []
|
||||
for msg_reaction in rankings:
|
||||
sorted_messages.append(msg_reaction.payload.payload.ranked_message_ids)
|
||||
logger.debug(f"SORTED MESSAGE {sorted_messages}")
|
||||
consensus = ranked_pairs(sorted_messages)
|
||||
ordered_ids_list: list[list[UUID]] = [
|
||||
msg_reaction.payload.payload.ranked_message_ids for msg_reaction in rankings
|
||||
]
|
||||
|
||||
common_set: set[UUID] = set.intersection(*map(set, ordered_ids_list))
|
||||
if len(common_set) < 2:
|
||||
logger.warning("The intersection of ranking results ID sets has less than two elements. Skipping.")
|
||||
continue
|
||||
|
||||
# keep only elements in commond set
|
||||
ordered_ids_list = [list(filter(lambda x: x in common_set, ids)) for ids in ordered_ids_list]
|
||||
assert all(len(x) == len(common_set) for x in ordered_ids_list)
|
||||
|
||||
logger.debug(f"SORTED MESSAGE IDS {ordered_ids_list}")
|
||||
consensus = ranked_pairs(ordered_ids_list)
|
||||
assert len(consensus) == len(common_set)
|
||||
logger.debug(f"CONSENSUS: {consensus}\n\n")
|
||||
|
||||
# fetch all siblings and clear ranks
|
||||
siblings = self.pr.fetch_message_siblings(consensus[0], reviewed=None, deleted=None)
|
||||
for m in siblings:
|
||||
m.rank = None
|
||||
self.db.add(m)
|
||||
|
||||
# index by id
|
||||
siblings = {m.id: m for m in siblings}
|
||||
|
||||
# set rank for each message that was part of the common set
|
||||
for rank, message_id in enumerate(consensus):
|
||||
# set rank for each message_id for Message rows
|
||||
msg = self.pr.fetch_message(message_id=message_id, fail_if_missing=True)
|
||||
msg.rank = rank
|
||||
self.db.add(msg)
|
||||
msg = siblings.get(message_id)
|
||||
if msg:
|
||||
msg.rank = rank
|
||||
self.db.add(msg)
|
||||
else:
|
||||
logger.warning(f"Message {message_id=} not found among siblings.")
|
||||
|
||||
except Exception:
|
||||
logger.exception(f"update_message_ranks({message_tree_id=}) failed")
|
||||
@@ -683,57 +765,65 @@ class TreeManager:
|
||||
# calculate acceptance based on spam label
|
||||
return np.mean([1 - l.labels[protocol_schema.TextLabel.spam] for l in labels])
|
||||
|
||||
def query_prompts_need_review(self, lang: str) -> list[Message]:
|
||||
"""
|
||||
Select initial prompt messages with less then required rankings in active message tree
|
||||
(active == True in message_tree_state)
|
||||
"""
|
||||
def _query_need_review(
|
||||
self, state: message_tree_state.State, required_reviews: int, root: bool, lang: str
|
||||
) -> list[Message]:
|
||||
|
||||
qry = (
|
||||
need_review = (
|
||||
self.db.query(Message)
|
||||
.select_from(MessageTreeState)
|
||||
.join(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
|
||||
.filter(
|
||||
MessageTreeState.active,
|
||||
MessageTreeState.state == message_tree_state.State.INITIAL_PROMPT_REVIEW,
|
||||
MessageTreeState.state == state,
|
||||
not_(Message.review_result),
|
||||
not_(Message.deleted),
|
||||
Message.review_count < self.cfg.num_reviews_initial_prompt,
|
||||
Message.parent_id.is_(None),
|
||||
Message.review_count < required_reviews,
|
||||
Message.lang == lang,
|
||||
)
|
||||
)
|
||||
|
||||
if root:
|
||||
need_review = need_review.filter(Message.parent_id.is_(None))
|
||||
else:
|
||||
need_review = need_review.filter(Message.parent_id.is_not(None))
|
||||
|
||||
if not settings.DEBUG_ALLOW_SELF_LABELING:
|
||||
qry = qry.filter(Message.user_id != self.pr.user_id)
|
||||
need_review = need_review.filter(Message.user_id != self.pr.user_id)
|
||||
|
||||
if settings.DEBUG_ALLOW_DUPLICATE_TASKS:
|
||||
qry = need_review
|
||||
else:
|
||||
user_id = self.pr.user_id
|
||||
need_review = need_review.cte(name="need_review")
|
||||
qry = (
|
||||
self.db.query(Message)
|
||||
.select_entity_from(need_review)
|
||||
.outerjoin(TextLabels, need_review.c.id == TextLabels.message_id)
|
||||
.group_by(need_review)
|
||||
.having(
|
||||
func.count(TextLabels.id).filter(TextLabels.task_id.is_not(None), TextLabels.user_id == user_id)
|
||||
== 0
|
||||
)
|
||||
)
|
||||
|
||||
return qry.all()
|
||||
|
||||
def query_prompts_need_review(self, lang: str) -> list[Message]:
|
||||
"""
|
||||
Select initial prompt messages with less then required rankings in active message tree
|
||||
(active == True in message_tree_state)
|
||||
"""
|
||||
return self._query_need_review(
|
||||
message_tree_state.State.INITIAL_PROMPT_REVIEW, self.cfg.num_reviews_initial_prompt, True, lang
|
||||
)
|
||||
|
||||
def query_replies_need_review(self, lang: str) -> list[Message]:
|
||||
"""
|
||||
Select child messages (parent_id IS NOT NULL) with less then required rankings
|
||||
in active message tree (active == True in message_tree_state)
|
||||
"""
|
||||
|
||||
qry = (
|
||||
self.db.query(Message)
|
||||
.select_from(MessageTreeState)
|
||||
.join(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
|
||||
.filter(
|
||||
MessageTreeState.active,
|
||||
MessageTreeState.state == message_tree_state.State.GROWING,
|
||||
not_(Message.review_result),
|
||||
not_(Message.deleted),
|
||||
Message.review_count < self.cfg.num_reviews_reply,
|
||||
Message.parent_id.is_not(None),
|
||||
Message.lang == lang,
|
||||
)
|
||||
)
|
||||
|
||||
if not settings.DEBUG_ALLOW_SELF_LABELING:
|
||||
qry = qry.filter(Message.user_id != self.pr.user_id)
|
||||
|
||||
return qry.all()
|
||||
return self._query_need_review(message_tree_state.State.GROWING, self.cfg.num_reviews_reply, False, lang)
|
||||
|
||||
_sql_find_incomplete_rankings = """
|
||||
-- find incomplete rankings
|
||||
@@ -749,17 +839,28 @@ WHERE mts.active -- only consider active trees
|
||||
AND m.parent_id IS NOT NULL -- ignore initial prompts
|
||||
GROUP BY m.parent_id, m.role
|
||||
HAVING COUNT(m.id) > 1 and MIN(m.ranking_count) < :num_required_rankings
|
||||
"""
|
||||
|
||||
_sql_find_incomplete_rankings_ex = f"""
|
||||
-- incomplete rankings but exclude of current user
|
||||
WITH incomplete_rankings AS ({_sql_find_incomplete_rankings})
|
||||
SELECT ir.* FROM incomplete_rankings ir
|
||||
LEFT JOIN message_reaction mr ON ir.parent_id = mr.message_id AND mr.payload_type = 'RankingReactionPayload'
|
||||
GROUP BY ir.parent_id, ir.role, ir.children_count, ir.child_min_ranking_count, ir.completed_rankings
|
||||
HAVING(COUNT(mr.message_id) FILTER (WHERE mr.user_id = :user_id) = 0)
|
||||
"""
|
||||
|
||||
def query_incomplete_rankings(self, lang: str) -> list[IncompleteRankingsRow]:
|
||||
"""Query parents which have childern that need further rankings"""
|
||||
|
||||
user_id = self.pr.user_id if not settings.DEBUG_ALLOW_DUPLICATE_TASKS else None
|
||||
r = self.db.execute(
|
||||
text(self._sql_find_incomplete_rankings),
|
||||
text(self._sql_find_incomplete_rankings_ex),
|
||||
{
|
||||
"num_required_rankings": self.cfg.num_required_rankings,
|
||||
"ranking_state": message_tree_state.State.RANKING,
|
||||
"lang": lang,
|
||||
"user_id": user_id,
|
||||
},
|
||||
)
|
||||
return [IncompleteRankingsRow.from_orm(x) for x in r.all()]
|
||||
@@ -780,17 +881,20 @@ WHERE mts.active -- only consider active trees
|
||||
AND (c.review_result OR coalesce(c.review_count, 0) < :num_reviews_reply) -- don't count children with negative review but count elements under review
|
||||
GROUP BY m.id, m.role, m.depth, m.message_tree_id, mts.max_children_count
|
||||
HAVING COUNT(c.id) < mts.max_children_count -- below maximum number of children
|
||||
AND COUNT(c.id) FILTER (WHERE c.user_id = :user_id) = 0 -- without reply by user
|
||||
"""
|
||||
|
||||
def query_extendible_parents(self, lang: str) -> list[ExtendibleParentRow]:
|
||||
"""Query parent messages that have not reached the maximum number of replies."""
|
||||
|
||||
user_id = self.pr.user_id if not settings.DEBUG_ALLOW_DUPLICATE_TASKS else None
|
||||
r = self.db.execute(
|
||||
text(self._sql_find_extendible_parents),
|
||||
{
|
||||
"growing_state": message_tree_state.State.GROWING,
|
||||
"num_reviews_reply": self.cfg.num_reviews_reply,
|
||||
"lang": lang,
|
||||
"user_id": user_id,
|
||||
},
|
||||
)
|
||||
return [ExtendibleParentRow.from_orm(x) for x in r.all()]
|
||||
@@ -814,12 +918,14 @@ HAVING COUNT(m.id) < mts.goal_tree_size
|
||||
def query_extendible_trees(self, lang: str) -> list[ActiveTreeSizeRow]:
|
||||
"""Query size of active message trees in growing state."""
|
||||
|
||||
user_id = self.pr.user_id if not settings.DEBUG_ALLOW_DUPLICATE_TASKS else None
|
||||
r = self.db.execute(
|
||||
text(self._sql_find_extendible_trees),
|
||||
{
|
||||
"growing_state": message_tree_state.State.GROWING,
|
||||
"num_reviews_reply": self.cfg.num_reviews_reply,
|
||||
"lang": lang,
|
||||
"user_id": user_id,
|
||||
},
|
||||
)
|
||||
return [ActiveTreeSizeRow.from_orm(x) for x in r.all()]
|
||||
@@ -827,18 +933,21 @@ HAVING COUNT(m.id) < mts.goal_tree_size
|
||||
def query_tree_size(self, message_tree_id: UUID) -> ActiveTreeSizeRow:
|
||||
"""Returns the number of reviewed not deleted messages in the message tree."""
|
||||
|
||||
required_reviews = settings.tree_manager.num_reviews_reply
|
||||
qry = (
|
||||
self.db.query(
|
||||
MessageTreeState.message_tree_id.label("message_tree_id"),
|
||||
MessageTreeState.goal_tree_size.label("goal_tree_size"),
|
||||
func.count(Message.id).label("tree_size"),
|
||||
func.count(Message.id).filter(Message.review_result).label("tree_size"),
|
||||
func.count(Message.id)
|
||||
.filter(not_(Message.review_result), Message.review_count < required_reviews)
|
||||
.label("awaiting_review"),
|
||||
)
|
||||
.select_from(MessageTreeState)
|
||||
.outerjoin(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
|
||||
.filter(
|
||||
MessageTreeState.active,
|
||||
not_(Message.deleted),
|
||||
Message.review_result,
|
||||
MessageTreeState.message_tree_id == message_tree_id,
|
||||
)
|
||||
.group_by(MessageTreeState.message_tree_id, MessageTreeState.goal_tree_size)
|
||||
@@ -907,7 +1016,7 @@ INNER JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'Ranki
|
||||
return rankings_by_message
|
||||
|
||||
@managed_tx_method(CommitMode.COMMIT)
|
||||
def ensure_tree_states(self):
|
||||
def ensure_tree_states(self) -> None:
|
||||
"""Add message tree state rows for all root nodes (inital prompt messages)."""
|
||||
|
||||
missing_tree_ids = self.query_misssing_tree_states()
|
||||
@@ -919,12 +1028,23 @@ INNER JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'Ranki
|
||||
logger.info(f"Inserting missing message tree state for message: {id} ({tree_size=}, {state=:s})")
|
||||
self._insert_default_state(id, state=state)
|
||||
|
||||
def query_num_active_trees(self, lang: str) -> int:
|
||||
rankings = (
|
||||
self.db.query(MessageTreeState).filter(MessageTreeState.state == message_tree_state.State.RANKING).all()
|
||||
)
|
||||
if len(rankings) > 0:
|
||||
logger.info(f"Checking state of {len(rankings)} message trees in ranking state.")
|
||||
for r in rankings:
|
||||
self.check_condition_for_scoring_state(r.message_tree_id)
|
||||
|
||||
def query_num_active_trees(self, lang: str, exclude_ranking: bool = True) -> int:
|
||||
"""Count all active trees (optionally exclude those in ranking state)."""
|
||||
query = (
|
||||
self.db.query(func.count(MessageTreeState.message_tree_id))
|
||||
.join(Message, MessageTreeState.message_tree_id == Message.id)
|
||||
.filter(MessageTreeState.active, Message.lang == lang)
|
||||
)
|
||||
if exclude_ranking:
|
||||
query = query.filter(MessageTreeState.state != message_tree_state.State.RANKING)
|
||||
return query.scalar()
|
||||
|
||||
def query_reviews_for_message(self, message_id: UUID) -> list[TextLabels]:
|
||||
@@ -1150,6 +1270,7 @@ DELETE FROM message WHERE message_tree_id = :message_tree_id;
|
||||
sql_purge_user = """
|
||||
DELETE FROM journal WHERE user_id = :user_id;
|
||||
DELETE FROM message_reaction WHERE user_id = :user_id;
|
||||
DELETE FROM message_emoji WHERE user_id = :user_id;
|
||||
DELETE FROM task WHERE user_id = :user_id;
|
||||
DELETE FROM message WHERE user_id = :user_id;
|
||||
DELETE FROM user_stats WHERE user_id = :user_id;
|
||||
@@ -1161,41 +1282,92 @@ DELETE FROM user_stats WHERE user_id = :user_id;
|
||||
if ban:
|
||||
self.db.execute(update(User).filter(User.id == user_id).values(deleted=True, enabled=False))
|
||||
|
||||
def export_trees_to_file(
|
||||
self,
|
||||
message_tree_ids: list[str],
|
||||
file=None,
|
||||
reviewed: bool = True,
|
||||
include_deleted: bool = False,
|
||||
use_compression: bool = False,
|
||||
) -> None:
|
||||
trees_to_export: List[tree_export.ExportMessageTree] = []
|
||||
|
||||
for message_tree_id in message_tree_ids:
|
||||
messages: List[Message] = self.pr.fetch_message_tree(message_tree_id, reviewed, include_deleted)
|
||||
trees_to_export.append(tree_export.build_export_tree(message_tree_id, messages))
|
||||
|
||||
if file:
|
||||
tree_export.write_trees_to_file(file, trees_to_export, use_compression)
|
||||
else:
|
||||
sys.stdout.write(json.dumps(jsonable_encoder(trees_to_export), indent=2))
|
||||
|
||||
def export_all_ready_trees(
|
||||
self, file: str, reviewed: bool = True, include_deleted: bool = False, use_compression: bool = False
|
||||
) -> None:
|
||||
message_tree_states: MessageTreeState = self.pr.fetch_message_trees_ready_for_export()
|
||||
message_tree_ids = [ms.message_tree_id for ms in message_tree_states]
|
||||
self.export_trees_to_file(message_tree_ids, file, reviewed, include_deleted, use_compression)
|
||||
|
||||
def export_all_user_trees(
|
||||
self,
|
||||
user_id: str,
|
||||
file: str,
|
||||
reviewed: bool = True,
|
||||
include_deleted: bool = False,
|
||||
use_compression: bool = False,
|
||||
) -> None:
|
||||
messages = self.pr.fetch_user_message_trees(UUID(user_id))
|
||||
message_tree_ids = [ms.message_tree_id for ms in messages]
|
||||
self.export_trees_to_file(message_tree_ids, file, reviewed, include_deleted, use_compression)
|
||||
|
||||
@managed_tx_method(CommitMode.COMMIT)
|
||||
def retry_scoring_failed_message_trees(self):
|
||||
query = self.db.query(MessageTreeState.message_tree_id).filter(
|
||||
MessageTreeState.state == message_tree_state.State.SCORING_FAILED
|
||||
)
|
||||
ranking_role_filter = None if self.cfg.rank_prompter_replies else "assistant"
|
||||
for row in query.all():
|
||||
try:
|
||||
message_tree_id = row["message_tree_id"]
|
||||
rankings_by_message = self.query_tree_ranking_results(message_tree_id, role_filter=ranking_role_filter)
|
||||
self.update_message_ranks(message_tree_id=message_tree_id, rankings_by_message=rankings_by_message)
|
||||
except Exception:
|
||||
logger.exception(f"retry_scoring_failed_message_trees failed for ({message_tree_id=})")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from oasst_backend.api.deps import api_auth
|
||||
|
||||
# from oasst_backend.api.deps import create_api_client
|
||||
from oasst_backend.database import engine
|
||||
from oasst_backend.prompt_repository import PromptRepository
|
||||
|
||||
with Session(engine) as db:
|
||||
api_client = api_auth(settings.OFFICIAL_WEB_API_KEY, db=db)
|
||||
# api_client = create_api_client(session=db, description="test", frontend_type="bot")
|
||||
dummy_user = protocol_schema.User(id="__dummy_user__", display_name="Dummy User", auth_method="local")
|
||||
|
||||
pr = PromptRepository(db=db, api_client=api_client, client_user=dummy_user)
|
||||
|
||||
cfg = TreeManagerConfiguration()
|
||||
tm = TreeManager(db, pr, cfg)
|
||||
tm.ensure_tree_states()
|
||||
|
||||
tm.purge_user_messages(user_id=UUID("2ef9ad21-0dc5-442d-8750-6f7f1790723f"), purge_initial_prompts=False)
|
||||
# tm.purge_user_messages(user_id=UUID("2ef9ad21-0dc5-442d-8750-6f7f1790723f"), purge_initial_prompts=False)
|
||||
# tm.purge_user(user_id=UUID("2ef9ad21-0dc5-442d-8750-6f7f1790723f"))
|
||||
# db.commit()
|
||||
|
||||
# print("query_num_active_trees", tm.query_num_active_trees())
|
||||
# print("query_incomplete_rankings", tm.query_incomplete_rankings())
|
||||
# print("query_replies_need_review", tm.query_replies_need_review())
|
||||
# print("query_incomplete_reply_reviews", tm.query_replies_need_review())
|
||||
# print("query_incomplete_initial_prompt_reviews", tm.query_prompts_need_review())
|
||||
# print("query_extendible_trees", tm.query_extendible_trees())
|
||||
# print("query_extendible_parents", tm.query_extendible_parents())
|
||||
# print("query_tree_size", tm.query_tree_size(message_tree_id=UUID("bdf434cf-4df5-4b74-949c-a5a157bc3292")))
|
||||
|
||||
# print(
|
||||
# "query_reviews_for_message",
|
||||
# tm.query_reviews_for_message(message_id=UUID("6a444493-0d48-4316-a9f1-7e263f5a2473")),
|
||||
# )
|
||||
|
||||
# print("next_task:", tm.next_task())
|
||||
|
||||
# print(
|
||||
# "query_tree_ranking_results", tm.query_tree_ranking_results(UUID("6036f58f-41b5-48c4-bdd9-b16f34ab1312"))
|
||||
# ".query_tree_ranking_results", tm.query_tree_ranking_results(UUID("2ac20d38-6650-43aa-8bb3-f61080c0d921"))
|
||||
# )
|
||||
|
||||
# print(tm.export_trees_to_file(message_tree_ids=["7e75fb38-e664-4e2b-817c-b9a0b01b0074"], file="lol.jsonl"))
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from oasst_backend.config import settings
|
||||
from oasst_backend.models import ApiClient, User
|
||||
from oasst_backend.utils.database_utils import CommitMode, managed_tx_method
|
||||
from oasst_shared.exceptions import OasstError, OasstErrorCode
|
||||
from oasst_shared.schemas import protocol as protocol_schema
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlmodel import Session, and_, or_
|
||||
from starlette.status import HTTP_403_FORBIDDEN, HTTP_404_NOT_FOUND
|
||||
|
||||
@@ -64,7 +66,13 @@ class UserRepository:
|
||||
return user
|
||||
|
||||
@managed_tx_method(CommitMode.COMMIT)
|
||||
def update_user(self, id: UUID, enabled: Optional[bool] = None, notes: Optional[str] = None) -> None:
|
||||
def update_user(
|
||||
self,
|
||||
id: UUID,
|
||||
enabled: Optional[bool] = None,
|
||||
notes: Optional[str] = None,
|
||||
show_on_leaderboard: Optional[bool] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Update a user by global user ID to disable or set admin notes. Only trusted clients may update users.
|
||||
|
||||
@@ -83,6 +91,8 @@ class UserRepository:
|
||||
user.enabled = enabled
|
||||
if notes is not None:
|
||||
user.notes = notes
|
||||
if show_on_leaderboard is not None:
|
||||
user.show_on_leaderboard = show_on_leaderboard
|
||||
|
||||
self.db.add(user)
|
||||
|
||||
@@ -107,9 +117,7 @@ class UserRepository:
|
||||
self.db.add(user)
|
||||
|
||||
@managed_tx_method(CommitMode.COMMIT)
|
||||
def lookup_client_user(self, client_user: protocol_schema.User, create_missing: bool = True) -> Optional[User]:
|
||||
if not client_user:
|
||||
return None
|
||||
def _lookup_client_user_tx(self, client_user: protocol_schema.User, create_missing: bool = True) -> Optional[User]:
|
||||
user: User = (
|
||||
self.db.query(User)
|
||||
.filter(
|
||||
@@ -135,6 +143,18 @@ class UserRepository:
|
||||
self.db.add(user)
|
||||
return user
|
||||
|
||||
def lookup_client_user(self, client_user: protocol_schema.User, create_missing: bool = True) -> Optional[User]:
|
||||
if not client_user:
|
||||
return None
|
||||
num_retries = settings.DATABASE_MAX_TX_RETRY_COUNT
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
return self._lookup_client_user_tx(client_user, create_missing)
|
||||
except IntegrityError:
|
||||
# catch UniqueViolation exception, for concurrent requests due to conflicts in ix_user_username
|
||||
if i + 1 == num_retries:
|
||||
raise
|
||||
|
||||
def query_users_ordered_by_username(
|
||||
self,
|
||||
api_client_id: Optional[UUID] = None,
|
||||
@@ -145,6 +165,7 @@ class UserRepository:
|
||||
auth_method: Optional[str] = None,
|
||||
search_text: Optional[str] = None,
|
||||
limit: Optional[int] = 100,
|
||||
desc: bool = False,
|
||||
) -> list[User]:
|
||||
if not self.api_client.trusted:
|
||||
if not api_client_id:
|
||||
@@ -184,14 +205,13 @@ class UserRepository:
|
||||
pattern = "%{}%".format(search_text.replace("\\", "\\\\").replace("_", "\\_").replace("%", "\\%"))
|
||||
qry = qry.filter(User.username.like(pattern))
|
||||
|
||||
if limit is not None and lte_username and not gte_username:
|
||||
# select top rows but return results in ascernding order
|
||||
sub_qry = qry.order_by(User.username.desc(), User.id.desc()).limit(limit).subquery("u")
|
||||
qry = self.db.query(User).select_entity_from(sub_qry).order_by(User.username, User.id)
|
||||
if desc:
|
||||
qry = qry.order_by(User.username.desc(), User.id.desc())
|
||||
else:
|
||||
qry = qry.order_by(User.username, User.id)
|
||||
if limit is not None:
|
||||
qry = qry.limit(limit)
|
||||
|
||||
if limit is not None:
|
||||
qry = qry.limit(limit)
|
||||
|
||||
return qry.all()
|
||||
|
||||
@@ -205,7 +225,9 @@ class UserRepository:
|
||||
auth_method: Optional[str] = None,
|
||||
search_text: Optional[str] = None,
|
||||
limit: Optional[int] = 100,
|
||||
desc: bool = False,
|
||||
) -> list[User]:
|
||||
|
||||
if not self.api_client.trusted:
|
||||
if not api_client_id:
|
||||
# Let unprivileged api clients query their own users without api_client_id being set
|
||||
@@ -255,13 +277,12 @@ class UserRepository:
|
||||
if auth_method:
|
||||
qry = qry.filter(User.auth_method == auth_method)
|
||||
|
||||
if limit is not None and lte_display_name and not gte_display_name:
|
||||
# select top rows but return results in ascernding order
|
||||
sub_qry = qry.order_by(User.display_name.desc(), User.id.desc()).limit(limit).subquery("u")
|
||||
qry = self.db.query(User).select_entity_from(sub_qry).order_by(User.display_name, User.id)
|
||||
if desc:
|
||||
qry = qry.order_by(User.display_name.desc(), User.id.desc())
|
||||
else:
|
||||
qry = qry.order_by(User.display_name, User.id)
|
||||
if limit is not None:
|
||||
qry = qry.limit(limit)
|
||||
|
||||
if limit is not None:
|
||||
qry = qry.limit(limit)
|
||||
|
||||
return qry.all()
|
||||
|
||||
@@ -39,7 +39,7 @@ class UserStatsRepository:
|
||||
qry = (
|
||||
self.session.query(User.id.label("user_id"), User.username, User.auth_method, User.display_name, UserStats)
|
||||
.join(UserStats, User.id == UserStats.user_id)
|
||||
.filter(UserStats.time_frame == time_frame.value)
|
||||
.filter(UserStats.time_frame == time_frame.value, User.show_on_leaderboard)
|
||||
.order_by(UserStats.rank)
|
||||
.limit(limit)
|
||||
)
|
||||
@@ -250,7 +250,8 @@ FROM
|
||||
PARTITION BY time_frame
|
||||
ORDER BY leader_score DESC, user_id
|
||||
) AS "rank", user_id, time_frame
|
||||
FROM user_stats
|
||||
FROM user_stats us2
|
||||
INNER JOIN "user" u ON us2.user_id = u.id AND u.show_on_leaderboard
|
||||
WHERE (:time_frame IS NULL OR time_frame = :time_frame)) AS r
|
||||
WHERE
|
||||
us.user_id = r.user_id
|
||||
|
||||
@@ -7,9 +7,14 @@ from loguru import logger
|
||||
from oasst_backend.config import settings
|
||||
from oasst_backend.database import engine
|
||||
from oasst_shared.exceptions import OasstError, OasstErrorCode
|
||||
from sqlalchemy.exc import OperationalError
|
||||
from psycopg2.errors import DeadlockDetected, ExclusionViolation, SerializationFailure, UniqueViolation
|
||||
from sqlalchemy.exc import OperationalError, PendingRollbackError
|
||||
from sqlmodel import Session, SQLModel
|
||||
|
||||
"""
|
||||
Error Handling Reference: https://www.postgresql.org/docs/15/mvcc-serialization-failure-handling.html
|
||||
"""
|
||||
|
||||
|
||||
class CommitMode(IntEnum):
|
||||
"""
|
||||
@@ -26,7 +31,6 @@ class CommitMode(IntEnum):
|
||||
* managed_tx_method and async_managed_tx_method methods are decorators functions
|
||||
* to be used on class functions. It expects the Class to have a 'db' Session object
|
||||
* initialised
|
||||
* TODO: tx method decorator for non class methods
|
||||
"""
|
||||
|
||||
|
||||
@@ -35,28 +39,46 @@ def managed_tx_method(auto_commit: CommitMode = CommitMode.COMMIT, num_retries=s
|
||||
@wraps(f)
|
||||
def wrapped_f(self, *args, **kwargs):
|
||||
try:
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
result = f(self, *args, **kwargs)
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
result = None
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
retry_exhausted = True
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
result = f(self, *args, **kwargs)
|
||||
self.db.commit()
|
||||
elif auto_commit == CommitMode.FLUSH:
|
||||
self.db.flush()
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
if isinstance(result, SQLModel):
|
||||
self.db.refresh(result)
|
||||
retry_exhausted = False
|
||||
break
|
||||
except PendingRollbackError as e:
|
||||
logger.info(str(e))
|
||||
self.db.rollback()
|
||||
except OperationalError as e:
|
||||
if e.orig is not None and isinstance(
|
||||
e.orig, (SerializationFailure, DeadlockDetected, UniqueViolation, ExclusionViolation)
|
||||
):
|
||||
logger.info(f"{type(e.orig)} Inner {e.orig.pgcode} {type(e.orig.pgcode)}")
|
||||
self.db.rollback()
|
||||
else:
|
||||
raise e
|
||||
logger.info(f"Retry {i+1}/{num_retries}")
|
||||
if retry_exhausted:
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
else:
|
||||
result = f(self, *args, **kwargs)
|
||||
if auto_commit == CommitMode.FLUSH:
|
||||
self.db.flush()
|
||||
if isinstance(result, SQLModel):
|
||||
self.db.refresh(result)
|
||||
return result
|
||||
except OperationalError:
|
||||
logger.info(f"Retry {i+1}/{num_retries} after possible DB concurrent update conflict.")
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
self.db.rollback()
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error("DB Rollback Failure")
|
||||
logger.info(str(e))
|
||||
raise e
|
||||
|
||||
return wrapped_f
|
||||
@@ -71,28 +93,46 @@ def async_managed_tx_method(
|
||||
@wraps(f)
|
||||
async def wrapped_f(self, *args, **kwargs):
|
||||
try:
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
result = await f(self, *args, **kwargs)
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
result = None
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
retry_exhausted = True
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
result = await f(self, *args, **kwargs)
|
||||
self.db.commit()
|
||||
elif auto_commit == CommitMode.FLUSH:
|
||||
self.db.flush()
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
if isinstance(result, SQLModel):
|
||||
self.db.refresh(result)
|
||||
retry_exhausted = False
|
||||
break
|
||||
except PendingRollbackError as e:
|
||||
logger.info(str(e))
|
||||
self.db.rollback()
|
||||
except OperationalError as e:
|
||||
if e.orig is not None and isinstance(
|
||||
e.orig, (SerializationFailure, DeadlockDetected, UniqueViolation, ExclusionViolation)
|
||||
):
|
||||
logger.info(f"{type(e.orig)} Inner {e.orig.pgcode} {type(e.orig.pgcode)}")
|
||||
self.db.rollback()
|
||||
else:
|
||||
raise e
|
||||
logger.info(f"Retry {i+1}/{num_retries}")
|
||||
if retry_exhausted:
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
else:
|
||||
result = await f(self, *args, **kwargs)
|
||||
if auto_commit == CommitMode.FLUSH:
|
||||
self.db.flush()
|
||||
if isinstance(result, SQLModel):
|
||||
self.db.refresh(result)
|
||||
return result
|
||||
except OperationalError:
|
||||
logger.info(f"Retry {i+1}/{num_retries} after possible DB concurrent update conflict.")
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
self.db.rollback()
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.exception("DB Rollback Failure")
|
||||
logger.info(str(e))
|
||||
raise e
|
||||
|
||||
return wrapped_f
|
||||
@@ -115,27 +155,49 @@ def managed_tx_function(
|
||||
@wraps(f)
|
||||
def wrapped_f(*args, **kwargs):
|
||||
try:
|
||||
for i in range(num_retries):
|
||||
with session_factory() as session:
|
||||
try:
|
||||
result = f(session, *args, **kwargs)
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
result = None
|
||||
if auto_commit == CommitMode.COMMIT:
|
||||
retry_exhausted = True
|
||||
for i in range(num_retries):
|
||||
with session_factory() as session:
|
||||
try:
|
||||
result = f(session, *args, **kwargs)
|
||||
session.commit()
|
||||
elif auto_commit == CommitMode.FLUSH:
|
||||
session.flush()
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
if isinstance(result, SQLModel):
|
||||
session.refresh(result)
|
||||
retry_exhausted = False
|
||||
break
|
||||
except PendingRollbackError as e:
|
||||
logger.info(str(e))
|
||||
session.rollback()
|
||||
return result
|
||||
except OperationalError:
|
||||
logger.info(f"Retry {i+1}/{num_retries} after possible DB concurrent update conflict.")
|
||||
session.rollback()
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
except OperationalError as e:
|
||||
if e.orig is not None and isinstance(
|
||||
e.orig,
|
||||
(SerializationFailure, DeadlockDetected, UniqueViolation, ExclusionViolation),
|
||||
):
|
||||
logger.info(f"{type(e.orig)} Inner {e.orig.pgcode} {type(e.orig.pgcode)}")
|
||||
session.rollback()
|
||||
else:
|
||||
raise e
|
||||
logger.info(f"Retry {i+1}/{num_retries}")
|
||||
if retry_exhausted:
|
||||
raise OasstError(
|
||||
"DATABASE_MAX_RETIRES_EXHAUSTED",
|
||||
error_code=OasstErrorCode.DATABASE_MAX_RETRIES_EXHAUSTED,
|
||||
http_status_code=HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
)
|
||||
else:
|
||||
with session_factory() as session:
|
||||
result = f(session, *args, **kwargs)
|
||||
if auto_commit == CommitMode.FLUSH:
|
||||
session.flush()
|
||||
if isinstance(result, SQLModel):
|
||||
session.refresh(result)
|
||||
elif auto_commit == CommitMode.ROLLBACK:
|
||||
session.rollback()
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error("DB Rollback Failure")
|
||||
logger.info(str(e))
|
||||
raise e
|
||||
|
||||
return wrapped_f
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
import os
|
||||
import pickle
|
||||
from collections import Counter
|
||||
|
||||
from sklearn import metrics
|
||||
from sklearn.feature_extraction.text import TfidfVectorizer
|
||||
from sklearn.model_selection import train_test_split
|
||||
from sklearn.pipeline import Pipeline
|
||||
from sklearn.svm import LinearSVC
|
||||
|
||||
|
||||
def load_and_split(foldername, num_words):
|
||||
ls = os.listdir(foldername)
|
||||
X = []
|
||||
Y = []
|
||||
langmap = dict()
|
||||
for idx, x in enumerate(ls):
|
||||
print("loading language", x)
|
||||
with open(foldername + "/" + x, "r") as reader:
|
||||
tmp = reader.read().split(" ")
|
||||
tmp = [" ".join(tmp[i : i + num_words]) for i in range(0, 100_000, num_words)]
|
||||
X.extend(tmp)
|
||||
Y.extend([idx] * len(tmp))
|
||||
langmap[idx] = x
|
||||
x_train, x_test, y_train, y_test = train_test_split(X, Y, test_size=0.90)
|
||||
return x_train, x_test, y_train, y_test, langmap
|
||||
|
||||
|
||||
def build_and_train_pipeline(x_train, y_train):
|
||||
vectorizer = TfidfVectorizer(ngram_range=(1, 2), analyzer="char", use_idf=False)
|
||||
clf = Pipeline(
|
||||
[
|
||||
("vec", vectorizer),
|
||||
# ("nystrom", Nystroem(n_components=1000,n_jobs=6)),
|
||||
("clf", LinearSVC(C=0.5)),
|
||||
# ("clf",GaussianNB())
|
||||
# ("clf", HistGradientBoostingClassifier())
|
||||
]
|
||||
)
|
||||
print("fitting model...")
|
||||
clf.fit(x_train, y_train)
|
||||
return clf
|
||||
|
||||
|
||||
def benchmark(clf, x_test, y_test, langmap):
|
||||
print("benchmarking model...")
|
||||
y_pred = clf.predict(x_test)
|
||||
names = list(langmap.values())
|
||||
# print(y_test)
|
||||
# print(langmap)
|
||||
print(metrics.classification_report(y_test, y_pred, target_names=names))
|
||||
cm = metrics.confusion_matrix(y_test, y_pred)
|
||||
print(cm)
|
||||
|
||||
|
||||
def main(foldername, modelname, num_words):
|
||||
x_train, x_test, y_train, y_test, langmap = load_and_split(foldername=foldername, num_words=num_words)
|
||||
clf = build_and_train_pipeline(x_train, y_train)
|
||||
benchmark(clf, x_test, y_test, langmap)
|
||||
save_model(clf, langmap, num_words, modelname)
|
||||
model = load(modelname)
|
||||
print(
|
||||
"running infernence on long tests",
|
||||
inference_voter(
|
||||
model,
|
||||
"""
|
||||
What language is this text written in? Nobody knows until you fill in at least ten words.
|
||||
This test here is to check whether the moving window approach works,
|
||||
so I still need to fill in a little more text.
|
||||
""",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def load(modelname):
|
||||
with open(modelname, "rb") as writer:
|
||||
data = pickle.load(writer)
|
||||
return data
|
||||
|
||||
|
||||
def save_model(model, idx_to_name, num_words, modelname):
|
||||
out = {
|
||||
"model": model,
|
||||
"idx_to_name": idx_to_name,
|
||||
"num_words": num_words,
|
||||
}
|
||||
with open(modelname, "wb") as writer:
|
||||
pickle.dump(out, writer)
|
||||
|
||||
|
||||
def inference_voter(model, text):
|
||||
tmp = text.split()
|
||||
# print(len(tmp), tmp)
|
||||
tmp = [" ".join(tmp[i : i + model["num_words"]]) for i in range(0, len(tmp) - model["num_words"])]
|
||||
predictions = model["model"].predict(tmp)
|
||||
# print("integer predictions", predictions)
|
||||
# print("name predictions", *[model["idx_to_name"][n] for n in predictions])
|
||||
result = Counter(predictions).most_common(1)[0][0]
|
||||
return model["idx_to_name"][result]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("-m", "--model", help="save location for model and metadata")
|
||||
parser.add_argument("-d", "--data", help="specify the folder for data files")
|
||||
parser.add_argument("-n", "--num_words", help="number of words to use for statistics", type=int)
|
||||
args = parser.parse_args()
|
||||
# np.set_printoptions(threshold=np.inf)
|
||||
main(args.data, args.model, args.num_words)
|
||||
@@ -96,13 +96,15 @@ def ranked_pairs(ranks: List[List[int]]):
|
||||
"""
|
||||
tallies, names = head_to_head_votes(ranks)
|
||||
tallies = tallies - tallies.T
|
||||
# print(tallies)
|
||||
# note: the resulting tally matrix should be skew-symmetric
|
||||
# order by strength of victory (using tideman's original method, don't think it would make a difference for us)
|
||||
sorted_majorities = []
|
||||
for i in range(len(ranks[0])):
|
||||
for j in range(len(ranks[0])):
|
||||
if tallies[i, j] > 0:
|
||||
# you can never prefer yourself over yourself
|
||||
# we also have to pick one of the two choices,
|
||||
# if the preference is exactly zero...
|
||||
if tallies[i, j] >= 0 and i != j:
|
||||
sorted_majorities.append((i, j, tallies[i, j]))
|
||||
# we don't explicitly deal with tied majorities here
|
||||
sorted_majorities = np.array(sorted(sorted_majorities, key=lambda x: x[2], reverse=True))
|
||||
@@ -128,13 +130,36 @@ def ranked_pairs(ranks: List[List[int]]):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
ranks = (
|
||||
|
||||
ranks = """ (
|
||||
[("w", "x", "z", "y") for _ in range(1)]
|
||||
+ [("w", "y", "x", "z") for _ in range(2)]
|
||||
# + [("x","y","z","w") for _ in range(4)]
|
||||
+ [("x", "z", "w", "y") for _ in range(5)]
|
||||
+ [("y", "w", "x", "z") for _ in range(1)]
|
||||
# [("y","z","w","x") for _ in range(1000)]
|
||||
)
|
||||
)"""
|
||||
ranks = [
|
||||
[
|
||||
("c5181083-d3e9-41e7-a935-83fb9fa01488"),
|
||||
("dcf3d179-0f34-4c15-ae21-b8feb15e422d"),
|
||||
("d11705af-5575-43e5-b22e-08d155fbaa62"),
|
||||
],
|
||||
[
|
||||
("d11705af-5575-43e5-b22e-08d155fbaa62"),
|
||||
("c5181083-d3e9-41e7-a935-83fb9fa01488"),
|
||||
("dcf3d179-0f34-4c15-ae21-b8feb15e422d"),
|
||||
],
|
||||
[
|
||||
("dcf3d179-0f34-4c15-ae21-b8feb15e422d"),
|
||||
("c5181083-d3e9-41e7-a935-83fb9fa01488"),
|
||||
("d11705af-5575-43e5-b22e-08d155fbaa62"),
|
||||
],
|
||||
[
|
||||
("d11705af-5575-43e5-b22e-08d155fbaa62"),
|
||||
("c5181083-d3e9-41e7-a935-83fb9fa01488"),
|
||||
("dcf3d179-0f34-4c15-ae21-b8feb15e422d"),
|
||||
],
|
||||
]
|
||||
rp = ranked_pairs(ranks)
|
||||
print(rp)
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import gzip
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from typing import Optional, TextIO
|
||||
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from oasst_backend.models import Message
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class ExportMessageNode(BaseModel):
|
||||
message_id: str
|
||||
parent_id: Optional[str]
|
||||
text: Optional[str]
|
||||
role: str
|
||||
review_count: Optional[int]
|
||||
rank: Optional[int]
|
||||
replies: Optional[list[ExportMessageNode]]
|
||||
|
||||
@classmethod
|
||||
def prep_message_export(cls, message: Message) -> ExportMessageNode:
|
||||
return cls(
|
||||
message_id=str(message.id),
|
||||
parent_id=str(message.parent_id) if message.parent_id else None,
|
||||
text=str(message.payload.payload.text),
|
||||
role=message.role,
|
||||
review_count=message.review_count,
|
||||
rank=message.rank,
|
||||
)
|
||||
|
||||
|
||||
class ExportMessageTree(BaseModel):
|
||||
message_tree_id: str
|
||||
replies: Optional[ExportMessageNode]
|
||||
|
||||
|
||||
def build_export_tree(message_tree_id: str, messages: list[Message]) -> ExportMessageTree:
|
||||
export_tree = ExportMessageTree(message_tree_id=str(message_tree_id))
|
||||
export_tree_data = [ExportMessageNode.prep_message_export(m) for m in messages]
|
||||
|
||||
message_parents = defaultdict(list)
|
||||
for message in export_tree_data:
|
||||
message_parents[message.parent_id].append(message)
|
||||
|
||||
def build_tree(tree: dict, parent: Optional[str], messages: list[Message]):
|
||||
children = message_parents[parent]
|
||||
tree.replies = children
|
||||
|
||||
for idx, child in enumerate(tree.replies):
|
||||
build_tree(tree.replies[idx], child.message_id, messages)
|
||||
|
||||
build_tree(export_tree, None, export_tree_data)
|
||||
|
||||
return export_tree
|
||||
|
||||
|
||||
def write_trees_to_file(file, trees: list[ExportMessageTree], use_compression: bool = True) -> None:
|
||||
|
||||
out_buff: TextIO
|
||||
if use_compression:
|
||||
out_buff = gzip.open(file, "wt", encoding="UTF-8")
|
||||
else:
|
||||
out_buff = open(file, "wt", encoding="UTF-8")
|
||||
|
||||
with out_buff as f:
|
||||
for tree in trees:
|
||||
file_data = jsonable_encoder(tree)
|
||||
json.dump(file_data, f)
|
||||
f.write("\n")
|
||||
Reference in New Issue
Block a user