Merge branch 'main' into 766_admin_enhancement

This commit is contained in:
notmd
2023-01-20 23:08:59 +07:00
83 changed files with 2770 additions and 449 deletions
@@ -0,0 +1,47 @@
"""switch to timestamp with tz
Revision ID: 7f0a28a156f4
Revises: 0964ac95170d
Create Date: 2023-01-19 21:53:01.107137
"""
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "7f0a28a156f4"
down_revision = "0964ac95170d"
branch_labels = None
depends_on = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.alter_column(table_name="user_stats", column_name="modified_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="user_stats", column_name="base_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="journal_integration", column_name="last_run", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="message_embedding", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="message_reaction", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="message_toxicity", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="message", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="task", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="task", column_name="expiry_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="text_labels", column_name="created_date", type_=sa.DateTime(timezone=True))
op.alter_column(table_name="user", column_name="created_date", type_=sa.DateTime(timezone=True))
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.alter_column(table_name="user_stats", column_name="modified_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="user_stats", column_name="base_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="journal_integration", column_name="last_run", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="message_embedding", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="message_reaction", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="message_toxicity", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="message", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="task", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="task", column_name="expiry_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="text_labels", column_name="created_date", type_=sa.DateTime(timezone=False))
op.alter_column(table_name="user", column_name="created_date", type_=sa.DateTime(timezone=False))
# ### end Alembic commands ###
+133 -1
View File
@@ -1,7 +1,17 @@
from datetime import datetime
from uuid import UUID
import pydantic
from fastapi import APIRouter, Depends
from loguru import logger
from oasst_backend.api import deps
from oasst_backend.config import Settings, settings
from oasst_backend.models import ApiClient, User
from oasst_backend.prompt_repository import PromptRepository
from oasst_backend.tree_manager import TreeManager
from oasst_backend.utils.database_utils import CommitMode, managed_tx_function
from oasst_shared.schemas.protocol import SystemStats
from oasst_shared.utils import ScopeTimer, unaware_to_utc
router = APIRouter()
@@ -13,7 +23,7 @@ class CreateApiClientRequest(pydantic.BaseModel):
admin_email: str | None = None
@router.post("/api_client")
@router.post("/api_client", response_model=str)
async def create_api_client(
request: CreateApiClientRequest,
root_token: str = Depends(deps.get_root_token),
@@ -29,3 +39,125 @@ async def create_api_client(
)
logger.info(f"Created api_client with key {api_client.api_key}")
return api_client.api_key
@router.get("/backend_settings/full", response_model=Settings)
async def get_backend_settings_full(api_client: ApiClient = Depends(deps.get_trusted_api_client)) -> Settings:
logger.info(
f"Backend settings requested by trusted api_client {api_client.id} (admin_email: {api_client.admin_email}, frontend_type: {api_client.frontend_type})"
)
return settings
class PublicSettings(pydantic.BaseModel):
"""Subset of backend settings which can be retrieved by untrusted API clients."""
PROJECT_NAME: str
API_V1_STR: str
DEBUG_USE_SEED_DATA: bool
DEBUG_ALLOW_SELF_LABELING: bool
DEBUG_SKIP_EMBEDDING_COMPUTATION: bool
DEBUG_SKIP_TOXICITY_CALCULATION: bool
DEBUG_DATABASE_ECHO: bool
USER_STATS_INTERVAL_DAY: int
USER_STATS_INTERVAL_WEEK: int
USER_STATS_INTERVAL_MONTH: int
USER_STATS_INTERVAL_TOTAL: int
@router.get("/backend_settings/public", response_model=PublicSettings)
async def get_backend_settings_public(api_client: ApiClient = Depends(deps.get_api_client)) -> PublicSettings:
return PublicSettings(**settings.dict())
class PurgeResultModel(pydantic.BaseModel):
before: SystemStats
after: SystemStats
preview: bool
duration: float
@router.post("/purge_user/{user_id}", response_model=PurgeResultModel)
async def purge_user(
user_id: UUID,
preview: bool = True,
ban: bool = True,
api_client: ApiClient = Depends(deps.get_trusted_api_client),
) -> str:
assert api_client.trusted
@managed_tx_function(CommitMode.ROLLBACK if preview else CommitMode.COMMIT)
def purge_tx(session: deps.Session) -> tuple[User, SystemStats, SystemStats]:
pr = PromptRepository(session, api_client)
stats_before = pr.get_stats()
user = pr.user_repository.get_user(user_id)
tm = TreeManager(session, pr)
tm.purge_user(user_id=user_id, ban=ban)
session.expunge(user)
return user, stats_before, pr.get_stats()
timer = ScopeTimer()
user, before, after = purge_tx()
timer.stop()
if preview:
logger.info(
f"PURGE USER PREVIEW: '{user.display_name}' (id: {str(user_id)}; username: '{user.username}'; auth-method: '{user.auth_method}')"
)
else:
logger.warning(
f"PURGE USER: '{user.display_name}' (id: {str(user_id)}; username: '{user.username}'; auth-method: '{user.auth_method}')"
)
logger.info(f"{before=}; {after=}")
return PurgeResultModel(before=before, after=after, preview=preview, duration=timer.elapsed)
@router.post("/purge_user/{user_id}/messages", response_model=PurgeResultModel)
async def purge_user_messages(
user_id: UUID,
purge_initial_prompts: bool = False,
min_date: datetime = None,
max_date: datetime = None,
preview: bool = True,
api_client: ApiClient = Depends(deps.get_trusted_api_client),
) -> str:
assert api_client.trusted
min_date = unaware_to_utc(min_date)
max_date = unaware_to_utc(max_date)
@managed_tx_function(CommitMode.ROLLBACK if preview else CommitMode.COMMIT)
def purge_user_messages_tx(session: deps.Session):
pr = PromptRepository(session, api_client)
stats_before = pr.get_stats()
user = pr.user_repository.get_user(user_id)
tm = TreeManager(session, pr)
tm.purge_user_messages(
user_id, purge_initial_prompts=purge_initial_prompts, min_date=min_date, max_date=max_date
)
session.expunge(user)
return user, stats_before, pr.get_stats()
timer = ScopeTimer()
user, before, after = purge_user_messages_tx()
timer.stop()
if preview:
logger.info(
f"PURGE USER MESSAGES PREVIEW: '{user.display_name}' (id: {str(user_id)}; username: '{user.username}'; auth-method: '{user.auth_method}')"
)
else:
logger.warning(
f"PURGE USER MESSAGES: '{user.display_name}' (id: {str(user_id)}; username: '{user.username}'; auth-method: '{user.auth_method}')"
)
logger.info(f"{before=}; {after=}")
return PurgeResultModel(before=before, after=after, preview=preview, duration=timer.elapsed)
@@ -45,7 +45,7 @@ def get_tree_by_frontend_id(
"""
pr = PromptRepository(db, api_client)
message = pr.fetch_message_by_frontend_message_id(message_id)
tree = pr.fetch_message_tree(message.message_tree_id)
tree = pr.fetch_message_tree(message.message_tree_id, reviewed=False)
return utils.prepare_tree(tree, message.message_tree_id)
@@ -6,6 +6,7 @@ from oasst_backend.models import ApiClient
from oasst_backend.user_stats_repository import UserStatsRepository, UserStatsTimeFrame
from oasst_shared.schemas.protocol import LeaderboardStats
from sqlmodel import Session
from starlette.status import HTTP_204_NO_CONTENT
router = APIRouter()
@@ -19,3 +20,22 @@ def get_leaderboard(
) -> LeaderboardStats:
usr = UserStatsRepository(db)
return usr.get_leaderboard(time_frame, limit=max_count)
@router.post("/update/{time_frame}", response_model=None, status_code=HTTP_204_NO_CONTENT)
def update_leaderboard_time_frame(
time_frame: UserStatsTimeFrame,
api_client: ApiClient = Depends(deps.get_trusted_api_client),
db: Session = Depends(deps.get_db),
) -> LeaderboardStats:
usr = UserStatsRepository(db)
return usr.update_stats(time_frame=time_frame)
@router.post("/update", response_model=None, status_code=HTTP_204_NO_CONTENT)
def update_leaderboards_all(
api_client: ApiClient = Depends(deps.get_trusted_api_client),
db: Session = Depends(deps.get_db),
) -> LeaderboardStats:
usr = UserStatsRepository(db)
return usr.update_all_time_frames()
+5 -1
View File
@@ -7,6 +7,7 @@ from oasst_backend.api.v1 import utils
from oasst_backend.models import ApiClient
from oasst_backend.prompt_repository import PromptRepository
from oasst_shared.schemas import protocol
from oasst_shared.utils import unaware_to_utc
from sqlmodel import Session
from starlette.status import HTTP_204_NO_CONTENT
@@ -29,6 +30,9 @@ def query_messages(
"""
Query messages.
"""
start_date = unaware_to_utc(start_date)
end_date = unaware_to_utc(end_date)
pr = PromptRepository(db, api_client)
messages = pr.query_messages(
username=username,
@@ -78,7 +82,7 @@ def get_tree(
"""
pr = PromptRepository(db, api_client)
message = pr.fetch_message(message_id)
tree = pr.fetch_message_tree(message.message_tree_id)
tree = pr.fetch_message_tree(message.message_tree_id, reviewed=False)
return utils.prepare_tree(tree, message.message_tree_id)
+32
View File
@@ -2,6 +2,7 @@ from fastapi import APIRouter, Depends
from oasst_backend.api import deps
from oasst_backend.models import ApiClient
from oasst_backend.prompt_repository import PromptRepository
from oasst_backend.tree_manager import TreeManager, TreeManagerStats, TreeMessageCountStats
from oasst_shared.schemas import protocol
from sqlmodel import Session
@@ -15,3 +16,34 @@ def get_message_stats(
):
pr = PromptRepository(db, api_client)
return pr.get_stats()
@router.get("/tree_manager/state_counts", response_model=dict[str, int])
def get_tree_manager__state_counts(
db: Session = Depends(deps.get_db),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
):
pr = PromptRepository(db, api_client)
tm = TreeManager(db, pr)
return tm.tree_counts_by_state()
@router.get("/tree_manager/message_counts", response_model=list[TreeMessageCountStats])
def get_tree_manager__message_counts(
only_active: bool = True,
db: Session = Depends(deps.get_db),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
):
pr = PromptRepository(db, api_client)
tm = TreeManager(db, pr)
return tm.tree_message_count_stats(only_active=only_active)
@router.get("/tree_manager", response_model=TreeManagerStats)
def get_tree_manager__stats(
db: Session = Depends(deps.get_db),
api_client: ApiClient = Depends(deps.get_trusted_api_client),
):
pr = PromptRepository(db, api_client)
tm = TreeManager(db, pr)
return tm.stats()
+2
View File
@@ -36,6 +36,8 @@ def request_task(
try:
pr = PromptRepository(db, api_client, client_user=request.user)
pr.ensure_user_is_enabled()
tm = TreeManager(db, pr)
task, message_tree_id, parent_message_id = tm.next_task(request.type)
pr.task_repository.store_task(task, message_tree_id, parent_message_id, request.collective)
+1 -1
View File
@@ -50,6 +50,6 @@ class JournalIntegration(SQLModel, table=True):
)
description: str = Field(max_length=512, primary_key=True)
last_journal_id: Optional[UUID] = Field(foreign_key="journal.id", nullable=True)
last_run: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(), nullable=True))
last_run: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(timezone=True), nullable=True))
last_error: Optional[str] = Field(nullable=True)
next_run: Optional[datetime] = Field(nullable=True)
+3 -1
View File
@@ -30,7 +30,9 @@ class Message(SQLModel, table=True):
api_client_id: UUID = Field(nullable=False, foreign_key="api_client.id")
frontend_message_id: str = Field(max_length=200, nullable=False)
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp(), index=True)
sa_column=sa.Column(
sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp(), index=True
)
)
payload_type: str = Field(nullable=False, max_length=200)
payload: Optional[PayloadContainer] = Field(
@@ -17,5 +17,5 @@ class MessageEmbedding(SQLModel, table=True):
# In the case that the Message Embedding is created afterwards
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp())
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp())
)
@@ -19,7 +19,9 @@ class MessageReaction(SQLModel, table=True):
sa_column=sa.Column(pg.UUID(as_uuid=True), sa.ForeignKey("user.id"), nullable=False, primary_key=True)
)
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp(), index=True)
sa_column=sa.Column(
sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp(), index=True
)
)
payload_type: str = Field(nullable=False, max_length=200)
payload: PayloadContainer = Field(sa_column=sa.Column(payload_column_type(PayloadContainer), nullable=False))
@@ -20,5 +20,5 @@ class MessageToxicity(SQLModel, table=True):
# In the case that the Message Embedding is created afterwards
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp())
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp())
)
+2 -2
View File
@@ -20,9 +20,9 @@ class Task(SQLModel, table=True):
),
)
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp()),
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp()),
)
expiry_date: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(), nullable=True))
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)
payload_type: str = Field(nullable=False, max_length=200)
payload: PayloadContainer = Field(sa_column=sa.Column(payload_column_type(PayloadContainer), nullable=False))
+3 -1
View File
@@ -17,7 +17,9 @@ class TextLabels(SQLModel, table=True):
)
user_id: UUID = Field(sa_column=sa.Column(pg.UUID(as_uuid=True), sa.ForeignKey("user.id"), nullable=False))
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp(), index=True),
sa_column=sa.Column(
sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp(), index=True
),
)
api_client_id: UUID = Field(nullable=False, foreign_key="api_client.id")
text: str = Field(nullable=False, max_length=2**16)
+1 -1
View File
@@ -21,7 +21,7 @@ class User(SQLModel, table=True):
auth_method: str = Field(nullable=False, max_length=128, default="local")
display_name: str = Field(nullable=False, max_length=256)
created_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp())
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp())
)
api_client_id: UUID = Field(foreign_key="api_client.id")
enabled: bool = Field(sa_column=sa.Column(sa.Boolean, nullable=False, server_default=sa.true()))
+2 -2
View File
@@ -26,11 +26,11 @@ class UserStats(SQLModel, table=True):
user_id: Optional[UUID] = Field(
sa_column=sa.Column(pg.UUID(as_uuid=True), sa.ForeignKey("user.id"), primary_key=True)
)
base_date: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(), nullable=True))
base_date: Optional[datetime] = Field(sa_column=sa.Column(sa.DateTime(timezone=True), nullable=True))
leader_score: int = 0
modified_date: Optional[datetime] = Field(
sa_column=sa.Column(sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp())
sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False, server_default=sa.func.current_timestamp())
)
rank: int = Field(nullable=True)
+31 -5
View File
@@ -28,8 +28,7 @@ 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 sqlalchemy import update
from sqlmodel import Session, func
from sqlmodel import Session, func, not_, text, update
from starlette.status import HTTP_403_FORBIDDEN, HTTP_404_NOT_FOUND
@@ -53,6 +52,13 @@ class PromptRepository:
)
self.journal = JournalWriter(db, api_client, self.user)
def ensure_user_is_enabled(self):
if self.user is None or self.user_id is None:
raise OasstError("User required", OasstErrorCode.USER_NOT_SPECIFIED)
if self.user.deleted or not self.user.enabled:
raise OasstError("User account disabled", OasstErrorCode.USER_DISABLED)
def fetch_message_by_frontend_message_id(self, frontend_message_id: str, fail_if_missing: bool = True) -> Message:
validate_frontend_message_id(frontend_message_id)
message: Message = (
@@ -146,6 +152,8 @@ class PromptRepository:
review_result: bool = False,
check_tree_state: bool = True,
) -> Message:
self.ensure_user_is_enabled()
validate_frontend_message_id(frontend_message_id)
validate_frontend_message_id(user_frontend_message_id)
@@ -354,8 +362,7 @@ class PromptRepository:
@managed_tx_method(CommitMode.FLUSH)
def insert_reaction(self, task_id: UUID, payload: db_payload.ReactionPayload) -> MessageReaction:
if self.user_id is None:
raise OasstError("User required", OasstErrorCode.USER_NOT_SPECIFIED)
self.ensure_user_is_enabled()
container = PayloadContainer(payload=payload)
reaction = MessageReaction(
@@ -499,10 +506,14 @@ class PromptRepository:
messages = self.db.query(Message).filter(Message.parent_id.is_(None)).order_by(func.random()).limit(size).all()
return messages
def fetch_message_tree(self, message_tree_id: UUID, reviewed: bool = True):
def fetch_message_tree(
self, message_tree_id: UUID, reviewed: bool = True, include_deleted: bool = False
) -> list[Message]:
qry = self.db.query(Message).filter(Message.message_tree_id == message_tree_id)
if reviewed:
qry = qry.filter(Message.review_result)
if not include_deleted:
qry = qry.filter(not_(Message.deleted))
return qry.all()
def fetch_multiple_random_replies(self, max_size: int = 5, message_role: str = None):
@@ -702,6 +713,21 @@ class PromptRepository:
return messages.all()
def update_children_counts(self, message_tree_id: UUID):
sql_update_children_count = """
UPDATE message SET children_count = cc.children_count
FROM (
SELECT m.id, count(c.id) - COALESCE(SUM(c.deleted::int), 0) AS children_count
FROM message m
LEFT JOIN message c ON m.id = c.parent_id
WHERE m.message_tree_id = :message_tree_id
GROUP BY m.id
) AS cc
WHERE message.id = cc.id;
"""
r = self.db.execute(text(sql_update_children_count), {"message_tree_id": message_tree_id})
logger.debug(f"update_children_count({message_tree_id=}): {r.rowcount} rows.")
@managed_tx_method(CommitMode.COMMIT)
def mark_messages_deleted(self, messages: Message | UUID | list[Message | UUID], recursive: bool = True):
"""
+293 -42
View File
@@ -1,4 +1,5 @@
import random
from datetime import datetime
from enum import Enum
from http import HTTPStatus
from typing import Any, Dict, List, Optional, Tuple
@@ -9,14 +10,14 @@ import pydantic
from loguru import logger
from oasst_backend.api.v1.utils import prepare_conversation, prepare_conversation_message_list
from oasst_backend.config import TreeManagerConfiguration, settings
from oasst_backend.models import Message, MessageReaction, MessageTreeState, Task, TextLabels, message_tree_state
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.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
from oasst_shared.exceptions.oasst_api_error import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol as protocol_schema
from sqlalchemy.sql import text
from sqlmodel import Session, func, not_
from sqlmodel import Session, func, not_, text, update
class TaskType(Enum):
@@ -68,6 +69,25 @@ class IncompleteRankingsRow(pydantic.BaseModel):
orm_mode = True
class TreeMessageCountStats(pydantic.BaseModel):
message_tree_id: UUID
state: str
depth: int
oldest: datetime
youngest: datetime
count: int
goal_tree_size: int
@property
def completed(self) -> int:
return self.count / self.goal_tree_size
class TreeManagerStats(pydantic.BaseModel):
state_counts: dict[str, int]
message_counts: list[TreeMessageCountStats]
class TreeManager:
_all_text_labels = list(map(lambda x: x.value, protocol_schema.TextLabel))
@@ -130,7 +150,7 @@ class TreeManager:
def _determine_task_availability_internal(
self,
num_active_trees: int,
extensible_parents: list[ExtendibleParentRow],
extendible_parents: list[ExtendibleParentRow],
prompts_need_review: list[Message],
replies_need_review: list[Message],
incomplete_rankings: list[IncompleteRankingsRow],
@@ -141,17 +161,17 @@ class TreeManager:
task_count_by_type[protocol_schema.TaskRequestType.initial_prompt] = num_missing_prompts
task_count_by_type[protocol_schema.TaskRequestType.prompter_reply] = len(
list(filter(lambda x: x.parent_role == "assistant", extensible_parents))
list(filter(lambda x: x.parent_role == "assistant", extendible_parents))
)
task_count_by_type[protocol_schema.TaskRequestType.assistant_reply] = len(
list(filter(lambda x: x.parent_role == "prompter", extensible_parents))
list(filter(lambda x: x.parent_role == "prompter", extendible_parents))
)
task_count_by_type[protocol_schema.TaskRequestType.label_initial_prompt] = len(prompts_need_review)
task_count_by_type[protocol_schema.TaskRequestType.label_assistant_reply] = len(
list(filter(lambda m: m.role == "assistant", replies_need_review))
)
task_count_by_type[protocol_schema.TaskRequestType.prompter_reply] = len(
task_count_by_type[protocol_schema.TaskRequestType.label_prompter_reply] = len(
list(filter(lambda m: m.role == "prompter", replies_need_review))
)
@@ -171,15 +191,17 @@ class TreeManager:
return task_count_by_type
def determine_task_availability(self) -> dict[protocol_schema.TaskRequestType, int]:
self.pr.ensure_user_is_enabled()
num_active_trees = self.query_num_active_trees()
extensible_parents = self.query_extendible_parents()
extendible_parents = self.query_extendible_parents()
prompts_need_review = self.query_prompts_need_review()
replies_need_review = self.query_replies_need_review()
incomplete_rankings = self.query_incomplete_rankings()
return self._determine_task_availability_internal(
num_active_trees=num_active_trees,
extensible_parents=extensible_parents,
extendible_parents=extendible_parents,
prompts_need_review=prompts_need_review,
replies_need_review=replies_need_review,
incomplete_rankings=incomplete_rankings,
@@ -191,10 +213,12 @@ class TreeManager:
logger.debug("TreeManager.next_task()")
self.pr.ensure_user_is_enabled()
num_active_trees = self.query_num_active_trees()
prompts_need_review = self.query_prompts_need_review()
replies_need_review = self.query_replies_need_review()
extensible_parents = self.query_extendible_parents()
extendible_parents = self.query_extendible_parents()
incomplete_rankings = self.query_incomplete_rankings()
if not self.cfg.rank_prompter_replies:
@@ -224,7 +248,7 @@ class TreeManager:
else:
task_count_by_type = self._determine_task_availability_internal(
num_active_trees=num_active_trees,
extensible_parents=extensible_parents,
extendible_parents=extendible_parents,
prompts_need_review=prompts_need_review,
replies_need_review=replies_need_review,
incomplete_rankings=incomplete_rankings,
@@ -266,7 +290,7 @@ class TreeManager:
ranking_parent_id = random.choice(incomplete_rankings).parent_id
messages = self.pr.fetch_message_conversation(ranking_parent_id)
assert len(messages) > 1 and messages[-1].id == ranking_parent_id
assert len(messages) > 0 and messages[-1].id == ranking_parent_id
ranking_parent = messages[-1]
assert not ranking_parent.deleted and ranking_parent.review_result
conversation = prepare_conversation(messages)
@@ -356,12 +380,12 @@ class TreeManager:
case TaskType.REPLY:
# select a tree with missing replies
if task_role == TaskRole.PROMPTER:
extensible_parents = list(filter(lambda x: x.parent_role == "assistant", extensible_parents))
extendible_parents = list(filter(lambda x: x.parent_role == "assistant", extendible_parents))
elif task_role == TaskRole.ASSISTANT:
extensible_parents = list(filter(lambda x: x.parent_role == "prompter", extensible_parents))
extendible_parents = list(filter(lambda x: x.parent_role == "prompter", extendible_parents))
if len(extensible_parents) > 0:
random_parent = random.choice(extensible_parents)
if len(extendible_parents) > 0:
random_parent = random.choice(extendible_parents)
# fetch random conversation to extend
logger.debug(f"selected {random_parent=}")
@@ -424,6 +448,7 @@ class TreeManager:
@async_managed_tx_method(CommitMode.COMMIT)
async def handle_interaction(self, interaction: protocol_schema.AnyInteraction) -> protocol_schema.Task:
pr = self.pr
pr.ensure_user_is_enabled()
match type(interaction):
case protocol_schema.TextReplyToMessage:
logger.info(
@@ -488,7 +513,8 @@ class TreeManager:
_, task = pr.store_ranking(interaction)
self.check_condition_for_scoring_state(task.message_tree_id)
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)
case protocol_schema.TextLabels:
logger.info(
@@ -551,7 +577,6 @@ class TreeManager:
mts = self.pr.fetch_tree_state(message_tree_id)
self._enter_state(mts, message_tree_state.State.ABORTED_LOW_GRADE)
@managed_tx_method(CommitMode.COMMIT)
def check_condition_for_growing_state(self, message_tree_id: UUID) -> bool:
logger.debug(f"check_condition_for_growing_state({message_tree_id=})")
@@ -569,7 +594,6 @@ class TreeManager:
self._enter_state(mts, message_tree_state.State.GROWING)
return True
@managed_tx_method(CommitMode.COMMIT)
def check_condition_for_ranking_state(self, message_tree_id: UUID) -> bool:
logger.debug(f"check_condition_for_ranking_state({message_tree_id=})")
@@ -587,22 +611,54 @@ class TreeManager:
self._enter_state(mts, message_tree_state.State.RANKING)
return True
def check_condition_for_scoring_state(self, message_tree_id: UUID) -> bool:
def check_condition_for_scoring_state(
self, message_tree_id: UUID
) -> Tuple[bool, dict[UUID, list[MessageReaction]]]:
logger.debug(f"check_condition_for_scoring_state({message_tree_id=})")
mts: MessageTreeState
mts = self.db.query(MessageTreeState).filter(MessageTreeState.message_tree_id == message_tree_id).one()
mts = self.pr.fetch_tree_state(message_tree_id)
if not mts.active or mts.state != message_tree_state.State.RANKING:
logger.debug(f"False {mts.active=}, {mts.state=}")
return False
return False, None
ranking_role_filter = None if self.cfg.rank_prompter_replies else "assistant"
rankings_by_message = self.query_tree_ranking_results(message_tree_id, role_filter=ranking_role_filter)
for parent_msg_id, ranking in rankings_by_message.items():
if len(ranking) < self.cfg.num_required_rankings:
logger.debug(f"False {parent_msg_id=} {len(ranking)=}")
return False
return False, None
self._enter_state(mts, message_tree_state.State.READY_FOR_SCORING)
return True, rankings_by_message
def update_message_ranks(self, message_tree_id: UUID, rankings_by_message: Dict[int, int]) -> bool:
mts = self.pr.fetch_tree_state(message_tree_id)
# check state, allow retry if in SCORING_FAILED state
if mts.state not in (message_tree_state.State.READY_FOR_SCORING, message_tree_state.State.SCORING_FAILED):
logger.debug(f"False {mts.active=}, {mts.state=}")
return False
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)
logger.debug(f"CONSENSUS: {consensus}\n\n")
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)
except Exception:
logger.exception(f"update_message_ranks({message_tree_id=}) failed")
self._enter_state(mts, message_tree_state.State.SCORING_FAILED)
return False
self._enter_state(mts, message_tree_state.State.READY_FOR_EXPORT)
return True
def _calculate_acceptance(self, labels: list[TextLabels]):
@@ -618,7 +674,7 @@ class TreeManager:
qry = (
self.db.query(Message)
.select_from(MessageTreeState)
.outerjoin(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
.join(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
.filter(
MessageTreeState.active,
MessageTreeState.state == message_tree_state.State.INITIAL_PROMPT_REVIEW,
@@ -643,7 +699,7 @@ class TreeManager:
qry = (
self.db.query(Message)
.select_from(MessageTreeState)
.outerjoin(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
.join(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
.filter(
MessageTreeState.active,
MessageTreeState.state == message_tree_state.State.GROWING,
@@ -664,7 +720,7 @@ class TreeManager:
SELECT m.parent_id, m.role, COUNT(m.id) children_count, MIN(m.ranking_count) child_min_ranking_count,
COUNT(m.id) FILTER (WHERE m.ranking_count >= :num_required_rankings) as completed_rankings
FROM message_tree_state mts
LEFT JOIN message m ON mts.message_tree_id = m.message_tree_id
INNER JOIN message m ON mts.message_tree_id = m.message_tree_id
WHERE mts.active -- only consider active trees
AND mts.state = :ranking_state -- message tree must be in ranking state
AND m.review_result -- must be reviewed
@@ -690,15 +746,15 @@ HAVING COUNT(m.id) > 1 and MIN(m.ranking_count) < :num_required_rankings
-- find all extendible parent nodes
SELECT m.id as parent_id, m.role as parent_role, m.depth, m.message_tree_id, COUNT(c.id) active_children_count
FROM message_tree_state mts
LEFT JOIN message m ON mts.message_tree_id = m.message_tree_id -- all elements of message tree
INNER JOIN message m ON mts.message_tree_id = m.message_tree_id -- all elements of message tree
LEFT JOIN message c ON m.id = c.parent_id -- child nodes
WHERE mts.active -- only consider active trees
AND mts.state = :growing_state -- message tree must be growing
AND NOT m.deleted -- ignore deleted messages as parents
AND m.depth < mts.max_depth -- ignore leaf nodes as parents
AND m.review_result -- parent node must have positive review
AND NOT c.deleted -- don't count deleted children
AND (c.review_result OR c.review_count < :num_reviews_reply) -- don't count children with negative review but count elements under review
AND NOT coalesce(c.deleted, FALSE) -- don't count deleted children
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
"""
@@ -708,7 +764,10 @@ HAVING COUNT(c.id) < mts.max_children_count -- below maximum number of children
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},
{
"growing_state": message_tree_state.State.GROWING,
"num_reviews_reply": self.cfg.num_reviews_reply,
},
)
return [ExtendibleParentRow.from_orm(x) for x in r.all()]
@@ -717,8 +776,8 @@ HAVING COUNT(c.id) < mts.max_children_count -- below maximum number of children
SELECT m.message_tree_id, mts.goal_tree_size, COUNT(m.id) AS tree_size
FROM (
SELECT DISTINCT message_tree_id FROM ({_sql_find_extendible_parents}) extendible_parents
) trees LEFT JOIN message_tree_state mts ON trees.message_tree_id = mts.message_tree_id
LEFT JOIN message m ON mts.message_tree_id = m.message_tree_id
) trees INNER JOIN message_tree_state mts ON trees.message_tree_id = mts.message_tree_id
INNER JOIN message m ON mts.message_tree_id = m.message_tree_id
WHERE NOT m.deleted
AND (
m.parent_id IS NOT NULL AND (m.review_result OR m.review_count < :num_reviews_reply) -- children
@@ -766,7 +825,7 @@ HAVING COUNT(m.id) < mts.goal_tree_size
"""Find all initial prompt messages that have no associated message tree state"""
qry_missing_tree_states = (
self.db.query(Message.id)
.join(MessageTreeState, isouter=True)
.outerjoin(MessageTreeState, Message.message_tree_id == MessageTreeState.message_tree_id)
.filter(
Message.parent_id.is_(None),
Message.message_tree_id == Message.id,
@@ -783,7 +842,7 @@ SELECT p.parent_id, mr.* FROM
-- find parents with > 1 children
SELECT m.parent_id, m.message_tree_id, COUNT(m.id) children_count
FROM message_tree_state mts
LEFT JOIN message m ON mts.message_tree_id = m.message_tree_id
INNER JOIN message m ON mts.message_tree_id = m.message_tree_id
WHERE m.review_result -- must be reviewed
AND NOT m.deleted -- not deleted
AND m.parent_id IS NOT NULL -- ignore initial prompts
@@ -792,8 +851,8 @@ SELECT p.parent_id, mr.* FROM
GROUP BY m.parent_id, m.message_tree_id
HAVING COUNT(m.id) > 1
) as p
LEFT JOIN task t ON p.parent_id = t.parent_message_id AND t.done AND (t.payload_type = 'RankPrompterRepliesPayload' OR t.payload_type = 'RankAssistantRepliesPayload')
LEFT JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'RankingReactionPayload'
INNER JOIN task t ON p.parent_id = t.parent_message_id AND t.done AND (t.payload_type = 'RankPrompterRepliesPayload' OR t.payload_type = 'RankAssistantRepliesPayload')
INNER JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'RankingReactionPayload'
"""
def query_tree_ranking_results(
@@ -832,7 +891,7 @@ LEFT JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'Rankin
state = message_tree_state.State.INITIAL_PROMPT_REVIEW
if tree_size > 1:
state = message_tree_state.State.GROWING
logger.info(f"Inserting missing message tree state for message: {id} ({tree_size=}, {state=})")
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) -> int:
@@ -885,6 +944,194 @@ LEFT JOIN message_reaction mr ON mr.task_id = t.id AND mr.payload_type = 'Rankin
active=True,
)
def tree_counts_by_state(self) -> dict[str, int]:
qry = self.db.query(
MessageTreeState.state, func.count(MessageTreeState.message_tree_id).label("count")
).group_by(MessageTreeState.state)
return {x["state"]: x["count"] for x in qry}
def tree_message_count_stats(self, only_active: bool = True) -> list[TreeMessageCountStats]:
qry = (
self.db.query(
MessageTreeState.message_tree_id,
func.max(Message.depth).label("depth"),
func.min(Message.created_date).label("oldest"),
func.max(Message.created_date).label("youngest"),
func.count(Message.id).label("count"),
MessageTreeState.goal_tree_size,
MessageTreeState.state,
)
.select_from(MessageTreeState)
.join(Message, MessageTreeState.message_tree_id == Message.message_tree_id)
.filter(not_(Message.deleted))
.group_by(MessageTreeState.message_tree_id)
)
if only_active:
qry = qry.filter(MessageTreeState.active)
return [TreeMessageCountStats(**x) for x in qry]
def stats(self) -> TreeManagerStats:
return TreeManagerStats(
state_counts=self.tree_counts_by_state(),
message_counts=self.tree_message_count_stats(only_active=True),
)
def get_user_messages_by_tree(
self,
user_id: UUID,
min_date: datetime = None,
max_date: datetime = None,
) -> Tuple[dict[UUID, list[Message]], list[Message]]:
"""Returns a dict with replies by tree (excluding initial prompts) and list of initial prompts
associated with user_id."""
# query all messages of the user
qry = self.db.query(Message).filter(Message.user_id == user_id)
if min_date:
qry = qry.filter(Message.created_date >= min_date)
if max_date:
qry = qry.filter(Message.created_date <= max_date)
prompts: list[Message] = []
replies_by_tree: dict[UUID, list[Message]] = {}
# walk over result set and distinguish between initial prompts and replies
for m in qry:
m: Message
if m.message_tree_id == m.id:
prompts.append(m)
else:
message_list = replies_by_tree.get(m.message_tree_id)
if message_list is None:
message_list = [m]
replies_by_tree[m.message_tree_id] = message_list
else:
message_list.append(m)
return replies_by_tree, prompts
def _purge_message_internal(self, message_id: UUID) -> None:
"""This internal function deletes a single message. It does not take care of
descendants, children_count in parent etc."""
sql_purge_message = """
DELETE FROM journal j USING message m WHERE j.message_id = :message_id;
DELETE FROM message_embedding e WHERE e.message_id = :message_id;
DELETE FROM message_toxicity t WHERE t.message_id = :message_id;
DELETE FROM text_labels l WHERE l.message_id = :message_id;
-- delete all ranking results that contain message
DELETE FROM message_reaction r WHERE r.payload_type = 'RankingReactionPayload' AND r.task_id IN (
SELECT t.id FROM message m
JOIN task t ON m.parent_id = t.parent_message_id
WHERE m.id = :message_id);
-- delete task which inserted message
DELETE FROM task t using message m WHERE t.id = m.task_id AND m.id = :message_id;
DELETE FROM task t WHERE t.parent_message_id = :message_id;
DELETE FROM message WHERE id = :message_id;
"""
r = self.db.execute(text(sql_purge_message), {"message_id": message_id})
logger.debug(f"purge_message({message_id=}): {r.rowcount} rows.")
def purge_message_tree(self, message_tree_id: UUID) -> None:
sql_purge_message_tree = """
DELETE FROM journal j USING message m WHERE j.message_id = m.Id AND m.message_tree_id = :message_tree_id;
DELETE FROM message_embedding e USING message m WHERE e.message_id = m.Id AND m.message_tree_id = :message_tree_id;
DELETE FROM message_toxicity t USING message m WHERE t.message_id = m.Id AND m.message_tree_id = :message_tree_id;
DELETE FROM text_labels l USING message m WHERE l.message_id = m.Id AND m.message_tree_id = :message_tree_id;
DELETE FROM message_reaction r USING task t WHERE r.task_id = t.id AND t.message_tree_id = :message_tree_id;
DELETE FROM task t WHERE t.message_tree_id = :message_tree_id;
DELETE FROM message_tree_state WHERE message_tree_id = :message_tree_id;
DELETE FROM message WHERE message_tree_id = :message_tree_id;
"""
r = self.db.execute(text(sql_purge_message_tree), {"message_tree_id": message_tree_id})
logger.debug(f"purge_message_tree({message_tree_id=}) {r.rowcount} rows.")
@managed_tx_method(CommitMode.FLUSH)
def purge_user_messages(
self,
user_id: UUID,
purge_initial_prompts: bool = True,
min_date: datetime = None,
max_date: datetime = None,
):
# find all affected message trees
replies_by_tree, prompts = self.get_user_messages_by_tree(user_id, min_date, max_date)
total_messages = sum(len(x) for x in replies_by_tree.values())
logger.debug(f"found: {len(replies_by_tree)} trees; {len(prompts)} prompts; {total_messages} messages;")
# remove all trees based on inital prompts of the user
if purge_initial_prompts:
for p in prompts:
self.purge_message_tree(p.message_tree_id)
if p.message_tree_id in replies_by_tree:
del replies_by_tree[p.message_tree_id]
# patch all affected message trees
for tree_id, replies in replies_by_tree.items():
bad_parent_ids = set(m.id for m in replies)
logger.debug(f"patching tree {tree_id=}, {bad_parent_ids=}")
tree_messages = self.pr.fetch_message_tree(tree_id, reviewed=False, include_deleted=True)
logger.debug(f"{tree_id=}, {len(bad_parent_ids)=}, {len(tree_messages)=}")
by_id = {m.id: m for m in tree_messages}
def ancestor_ids(msg: Message) -> list[UUID]:
t = []
while msg.parent_id is not None:
msg = by_id[msg.parent_id]
t.append(msg.id)
return t
def is_descendant_of_deleted(m: Message) -> bool:
if m.id in bad_parent_ids:
return True
ancestors = ancestor_ids(m)
if any(a in bad_parent_ids for a in ancestors):
return True
return False
# start with deepest messages first
tree_messages.sort(key=lambda x: x.depth, reverse=True)
for m in tree_messages:
if is_descendant_of_deleted(m):
logger.debug(f"purging message: {m.id}")
self._purge_message_internal(m.id)
# update childern counts
self.pr.update_children_counts(m.message_tree_id)
# reactivate tree
logger.info(f"reactivating message tree {tree_id}")
mts = self.pr.fetch_tree_state(tree_id)
mts.active = True
self._enter_state(mts, message_tree_state.State.INITIAL_PROMPT_REVIEW)
self.check_condition_for_growing_state(tree_id)
self.check_condition_for_ranking_state(tree_id)
self.check_condition_for_scoring_state(tree_id)
@managed_tx_method(CommitMode.FLUSH)
def purge_user(self, user_id: UUID, ban: bool = True) -> None:
self.purge_user_messages(user_id, purge_initial_prompts=True)
# delete all remaining rows and ban user
sql_purge_user = """
DELETE FROM journal WHERE user_id = :user_id;
DELETE FROM message_reaction 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;
"""
r = self.db.execute(text(sql_purge_user), {"user_id": user_id})
logger.debug(f"purge_user({user_id=}): {r.rowcount} rows.")
if ban:
self.db.execute(update(User).filter(User.id == user_id).values(deleted=True, enabled=False))
if __name__ == "__main__":
from oasst_backend.api.deps import api_auth
@@ -901,6 +1148,10 @@ if __name__ == "__main__":
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(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())
@@ -909,10 +1160,10 @@ if __name__ == "__main__":
# 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(
# "query_reviews_for_message",
# tm.query_reviews_for_message(message_id=UUID("6a444493-0d48-4316-a9f1-7e263f5a2473")),
# )
# print("next_task:", tm.next_task())
+13 -10
View File
@@ -4,6 +4,7 @@ from uuid import UUID
import sqlalchemy as sa
from loguru import logger
from oasst_backend.config import settings
from oasst_backend.models import Message, MessageReaction, Task, User, UserStats, UserStatsTimeFrame
from oasst_backend.models.db_payload import (
LabelAssistantReplyPayload,
@@ -39,12 +40,16 @@ class UserStatsRepository:
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)
.order_by(UserStats.leader_score.desc())
.order_by(UserStats.rank)
.limit(limit)
)
leaderboard = [_create_user_score(r) for r in self.session.exec(qry)]
return LeaderboardStats(time_frame=time_frame.value, leaderboard=leaderboard)
if len(leaderboard) > 0:
last_update = max(x.modified_date for x in leaderboard)
else:
last_update = utcnow()
return LeaderboardStats(time_frame=time_frame.value, leaderboard=leaderboard, last_updated=last_update)
def get_user_stats_all_time_frames(self, user_id: UUID) -> dict[str, UserScore | None]:
qry = (
@@ -291,13 +296,11 @@ WHERE
if __name__ == "__main__":
from oasst_backend.api.deps import get_dummy_api_client
from oasst_backend.api.deps import api_auth
from oasst_backend.database import engine
with Session(engine) as session:
api_client = get_dummy_api_client(session)
usr = UserStatsRepository(session)
# usr.update_all_time_frames()
# session.commit()
# usr.get_leader_board(UserStatsTimeFrame.total)
usr.get_user_stats_all_time_frames(UUID("0d6ff62a-0bea-4c56-ade8-b3e0520a10ce"))
with Session(engine) as db:
api_client = api_auth(settings.OFFICIAL_WEB_API_KEY, db=db)
usr = UserStatsRepository(db)
usr.update_all_time_frames()
db.commit()
@@ -19,6 +19,7 @@ class CommitMode(IntEnum):
NONE = 0
FLUSH = 1
COMMIT = 2
ROLLBACK = 3
"""
@@ -41,6 +42,8 @@ def managed_tx_method(auto_commit: CommitMode = CommitMode.COMMIT, num_retries=s
self.db.commit()
elif auto_commit == CommitMode.FLUSH:
self.db.flush()
elif auto_commit == CommitMode.ROLLBACK:
self.db.rollback()
if isinstance(result, SQLModel):
self.db.refresh(result)
return result
@@ -75,6 +78,8 @@ def async_managed_tx_method(
self.db.commit()
elif auto_commit == CommitMode.FLUSH:
self.db.flush()
elif auto_commit == CommitMode.ROLLBACK:
self.db.rollback()
if isinstance(result, SQLModel):
self.db.refresh(result)
return result
@@ -118,8 +123,8 @@ def managed_tx_function(
session.commit()
elif auto_commit == CommitMode.FLUSH:
session.flush()
if isinstance(result, SQLModel):
session.refresh(result)
elif auto_commit == CommitMode.ROLLBACK:
session.rollback()
return result
except OperationalError:
logger.info(f"Retry {i+1}/{num_retries} after possible DB concurrent update conflict.")
+140
View File
@@ -0,0 +1,140 @@
from typing import List
import numpy as np
def head_to_head_votes(ranks: List[List[int]]):
tallies = np.zeros((len(ranks[0]), len(ranks[0])))
names = sorted(ranks[0])
ranks = np.array(ranks)
# we want the sorted indices
ranks = np.argsort(ranks, axis=1)
for i in range(ranks.shape[1]):
for j in range(i + 1, ranks.shape[1]):
# now count the cases someone voted for i over j
over_j = np.sum(ranks[:, i] < ranks[:, j])
over_i = np.sum(ranks[:, j] < ranks[:, i])
tallies[i, j] = over_j
# tallies[i,j] = over_i
tallies[j, i] = over_i
# tallies[j,i] = over_j
return tallies, names
def cycle_detect(pairs):
"""Recursively detect cylces by removing condorcet losers until either only one pair is left or condorcet loosers no longer exist
This method upholds the invariant that in a ranking for all a,b either a>b or b>a for all a,b.
Returns
-------
out : False if the pairs do not contain a cycle, True if the pairs contain a cycle
"""
# get all condorcet losers (pairs that loose to all other pairs)
# idea: filter all losers that are never winners
# print("pairs", pairs)
if len(pairs) <= 1:
return False
losers = [c_lose for c_lose in np.unique(pairs[:, 1]) if c_lose not in pairs[:, 0]]
if len(losers) == 0:
# if we recursively removed pairs, and at some point we did not have
# a condorcet loser, that means everything is both a winner and loser,
# yielding at least one (winner,loser), (loser,winner) pair
return True
new = []
for p in pairs:
if p[1] not in losers:
new.append(p)
return cycle_detect(np.array(new))
def get_winner(pairs):
"""
This returns _one_ concordant winner.
It could be that there are multiple concordant winners, but in our case
since we are interested in a ranking, we have to choose one at random.
"""
losers = np.unique(pairs[:, 1]).astype(int)
winners = np.unique(pairs[:, 0]).astype(int)
for w in winners:
if w not in losers:
return w
def get_ranking(pairs):
"""
Abuses concordance property to get a (not necessarily unqiue) ranking.
The lack of uniqueness is due to the potential existence of multiple
equally ranked winners. We have to pick one, which is where
the non-uniqueness comes from
"""
if len(pairs) == 1:
return list(pairs[0])
w = get_winner(pairs)
# now remove the winner from the list of pairs
p_new = np.array([(a, b) for a, b in pairs if a != w])
return [w] + get_ranking(p_new)
def ranked_pairs(ranks: List[List[int]]):
"""
Expects a list of rankings for an item like:
[("w","x","z","y") for _ in range(3)]
+ [("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)]
This code is quite brain melting, but the idea is the following:
1. create a head-to-head matrix that tallies up all win-lose combinations of preferences
2. take all combinations that win more than they loose and sort those by how often they win
3. use that to create an (implicit) directed graph
4. recursively extract nodes from the graph that do not have incoming edges
5. said recursive list is the ranking
"""
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:
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))
# now do lock ins
lock_ins = []
for (x, y, _) in sorted_majorities:
# invariant: lock_ins has no cycles here
lock_ins.append((x, y))
# print("lock ins are now",np.array(lock_ins))
if cycle_detect(np.array(lock_ins)):
# print("backup: cycle detected")
# if there's a cycle, delete the new addition and continue
lock_ins = lock_ins[:-1]
# now simply return all winners in order, and attach the losers
# to the back. This is because the overall loser might not be unique
# and (by concordance property) may never exist in any winning set to begin with.
# (otherwise he would either not be the loser, or cycles exist!)
# Since there could be multiple overall losers, we just return them in any order
# as we are unable to find a closer ranking
numerical_ranks = np.array(get_ranking(np.array(lock_ins))).astype(int)
conversion = [names[n] for n in numerical_ranks]
return conversion
if __name__ == "__main__":
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)]
)
rp = ranked_pairs(ranks)
print(rp)