mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-07-24 12:50:50 +08:00
Merge branch 'main' into 766_admin_enhancement
This commit is contained in:
@@ -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 ###
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
)
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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.")
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user