mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-07-28 11:15:43 +08:00
unified queueing
This commit is contained in:
@@ -7,8 +7,16 @@ from sqlalchemy.sql.operators import is_not
|
||||
|
||||
|
||||
class ChatRepository:
|
||||
def __init__(self, session: sqlmodel.Session) -> None:
|
||||
def __init__(self, session: sqlmodel.Session, do_commit=True) -> None:
|
||||
self.session = session
|
||||
self.do_commit = do_commit
|
||||
|
||||
def as_no_commit(self) -> "ChatRepository":
|
||||
return ChatRepository(self.session, do_commit=False)
|
||||
|
||||
def maybe_commit(self) -> None:
|
||||
if self.do_commit:
|
||||
self.session.commit()
|
||||
|
||||
def get_chats(self) -> list[models.DbChatEntry]:
|
||||
return self.session.exec(sqlmodel.select(models.DbChatEntry)).all()
|
||||
@@ -25,22 +33,25 @@ class ChatRepository:
|
||||
chats = self.get_chats()
|
||||
return [chat.to_list_entry() for chat in chats]
|
||||
|
||||
def get_chat_by_id(self, id: str) -> models.DbChatEntry:
|
||||
chat = self.session.exec(sqlmodel.select(models.DbChatEntry).where(models.DbChatEntry.id == id)).one()
|
||||
def get_chat_by_id(self, chat_id: str, for_update=False) -> models.DbChatEntry:
|
||||
query = sqlmodel.select(models.DbChatEntry).where(models.DbChatEntry.id == chat_id)
|
||||
if for_update:
|
||||
query = query.with_for_update()
|
||||
chat = self.session.exec(query).one()
|
||||
return chat
|
||||
|
||||
def get_chat_entry_by_id(self, id: str) -> interface.ChatEntry:
|
||||
return self.get_chat_by_id(id).to_entry()
|
||||
def get_chat_entry_by_id(self, chat_id: str) -> interface.ChatEntry:
|
||||
return self.get_chat_by_id(chat_id).to_entry()
|
||||
|
||||
def create_chat(self) -> models.DbChatEntry:
|
||||
chat = models.DbChatEntry()
|
||||
self.session.add(chat)
|
||||
self.session.commit()
|
||||
self.maybe_commit()
|
||||
return chat
|
||||
|
||||
def add_prompter_message(self, id: str, message_request: interface.MessageRequest) -> None:
|
||||
logger.info(f"Adding prompter message {message_request} to chat {id}")
|
||||
chat = self.get_chat_by_id(id)
|
||||
def add_prompter_message(self, chat_id: str, message_request: interface.MessageRequest) -> None:
|
||||
logger.info(f"Adding prompter message {message_request} to chat {chat_id}")
|
||||
chat = self.get_chat_by_id(chat_id, for_update=True)
|
||||
if not chat.conversation.is_prompter_turn:
|
||||
raise fastapi.HTTPException(status_code=400, detail="Not your turn")
|
||||
if chat.pending_message_request is not None:
|
||||
@@ -55,12 +66,12 @@ class ChatRepository:
|
||||
|
||||
chat.pending_message_request = message_request
|
||||
chat.message_request_state = interface.MessageRequestState.pending
|
||||
self.session.commit()
|
||||
logger.debug(f"Added prompter message {message_request} to chat {id}")
|
||||
self.maybe_commit()
|
||||
logger.debug(f"Added prompter message {message_request} to chat {chat_id}")
|
||||
|
||||
def add_assistant_message(self, id: str, text: str) -> None:
|
||||
logger.info(f"Adding assistant message {text} to chat {id}")
|
||||
chat = self.get_chat_by_id(id)
|
||||
def add_assistant_message(self, chat_id: str, text: str) -> None:
|
||||
logger.info(f"Adding assistant message {text} to chat {chat_id}")
|
||||
chat = self.get_chat_by_id(chat_id, for_update=True)
|
||||
chat.conversation.messages.append(
|
||||
protocol.ConversationMessage(
|
||||
text=text,
|
||||
@@ -68,12 +79,12 @@ class ChatRepository:
|
||||
)
|
||||
)
|
||||
chat.pending_message_request = None
|
||||
self.session.commit()
|
||||
logger.debug(f"Added assistant message {text} to chat {id}")
|
||||
self.maybe_commit()
|
||||
logger.debug(f"Added assistant message {text} to chat {chat_id}")
|
||||
|
||||
def set_chat_state(self, id: str, state: interface.MessageRequestState) -> None:
|
||||
logger.info(f"Setting chat {id} state to {state}")
|
||||
chat = self.get_chat_by_id(id)
|
||||
def set_chat_state(self, chat_id: str, state: interface.MessageRequestState) -> None:
|
||||
logger.info(f"Setting chat {chat_id} state to {state}")
|
||||
chat = self.get_chat_by_id(chat_id, for_update=True)
|
||||
chat.message_request_state = state
|
||||
self.session.commit()
|
||||
logger.debug(f"Set chat {id} state to {state}")
|
||||
self.maybe_commit()
|
||||
logger.debug(f"Set chat {chat_id} state to {state}")
|
||||
|
||||
@@ -9,8 +9,9 @@ class MessageRequest(pydantic.BaseModel):
|
||||
model_name: str = "distilgpt2"
|
||||
max_new_tokens: int = 100
|
||||
|
||||
def compatible_with(self, worker_config: inference.WorkerConfig) -> bool:
|
||||
return self.model_name == worker_config.model_name
|
||||
@property
|
||||
def worker_compat_hash(self) -> str:
|
||||
return f"{self.model_name}"
|
||||
|
||||
|
||||
class TokenResponseEvent(pydantic.BaseModel):
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
import redis.asyncio as redis
|
||||
from oasst_inference_server.settings import settings
|
||||
|
||||
|
||||
class RedisQueue:
|
||||
def __init__(self, redis_client: redis.Redis, queue_id: str) -> None:
|
||||
self.redis_client = redis_client
|
||||
self.queue_id = queue_id
|
||||
|
||||
async def enqueue(self, value: str) -> None:
|
||||
return await self.redis_client.rpush(self.queue_id, value)
|
||||
|
||||
async def dequeue(self, block: bool = True, timeout: int = 1) -> str:
|
||||
if block:
|
||||
return await self.redis_client.blpop(self.queue_id, timeout=timeout)
|
||||
else:
|
||||
return await self.redis_client.lpop(self.queue_id)
|
||||
|
||||
|
||||
def chat_queue(redis_client: redis.Redis, chat_id: str) -> RedisQueue:
|
||||
return RedisQueue(redis_client, f"chat:{chat_id}")
|
||||
|
||||
|
||||
def work_queue(redis_client: redis.Redis, worker_compat_hash: str) -> RedisQueue:
|
||||
if worker_compat_hash not in settings.allowed_worker_compat_hashes:
|
||||
raise ValueError(f"Worker compat hash {worker_compat_hash} not allowed")
|
||||
return RedisQueue(redis_client, f"work:{worker_compat_hash}")
|
||||
@@ -8,6 +8,8 @@ class Settings(pydantic.BaseSettings):
|
||||
redis_port: int = 6379
|
||||
redis_db: int = 0
|
||||
|
||||
allowed_worker_compat_hashes: list[str] = ["distilgpt2"]
|
||||
|
||||
sse_retry_timeout: int = 15000
|
||||
update_alembic: bool = True
|
||||
alembic_retries: int = 5
|
||||
|
||||
Reference in New Issue
Block a user