Merging from main

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