@@ -264,6 +264,8 @@ SSE_POLL_INTERVAL = env.float('SSE_POLL_INTERVAL', 0.2) # Celery CELERY_BROKER_URL = env.str('CELERY_BROKER_URL', 'redis://celery-mdb:6379/0') CELERY_RESULT_BACKEND = env.str('CELERY_RESULT_BACKEND', 'redis://celery-mdb:6379/0') +CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT = env.int('CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT', 630) +CELERY_WORKER_PREFETCH_MULTIPLIER = env.int('CELERY_WORKER_PREFETCH_MULTIPLIER', 1) CELERY_ACCEPT_CONTENT = ['json', 'application/x-python-serialize', 'pickle'] CELERY_RESULT_SERIALIZER = 'pickle' @@ -143,7 +143,7 @@ class NeuronModel(BaseModel, OrderedModel): @property def service(self) -> 'SimpleService': # noqa: F821 - return getattr(importlib.import_module(f'ml_model.services.{self.slug}'), f'{self.slug.title()}') + return getattr(importlib.import_module(f'ml_model.services.{self.slug}'), self.slug.replace('-', '').title()) @property def streaming(self) -> bool: @@ -30,9 +30,9 @@ def get_links(request): ) -@router.get('{chat_uid}/messages/{message_uid}/stream', tags=['chats']) -def stream_message_reconnect(request, chat_uid: UUID, message_uid: UUID, offset: int): - store = SSEStoreService(message_uuid=message_uid, user_uuid=request.auth.uid, chat_uuid=chat_uid) +@router.get('{chat_uid}/messages/stream', tags=['chats']) +def stream_message_reconnect(request, chat_uid: UUID, offset: int = 0): + store = SSEStoreService(user_uuid=request.auth.uid, chat_uuid=chat_uid) if not store.exists(): raise HttpError(404, _('Stream not found')) sse_chat_stream = SSEChatStreamService(store) @@ -53,12 +53,12 @@ def stream_message(request, chat_uid: UUID, body: MessageInSchema): if not chat.model.streaming: raise HttpError(501, _('Stream not supported for this model')) - data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) - input_message = Message.objects.create(content_object=chat, from_model=False, **data) - - store = SSEStoreService(message_uuid=input_message.pk, user_uuid=request.auth.uid, chat_uuid=chat.pk) + store = SSEStoreService(user_uuid=request.auth.uid, chat_uuid=chat_uid) if store.exists(): raise HttpError(409, _('Stream already in progress')) + + data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + input_message = Message.objects.create(content_object=chat, from_model=False, **data) store.start() event_stream_task.delay( @@ -19,6 +19,7 @@ class SSEChatStreamService: def event_stream(self, request: HttpRequest | None = None, offset: int = 0) -> Iterator[str]: last_event_id = offset + message_uuid = self.store.get_message_uuid() try: while self.store.exists(): if getattr(request, 'closed', False): @@ -31,6 +32,8 @@ class SSEChatStreamService: continue for chunk in chunks: + if chunk.event == 'start': + message_uuid = chunk.data.get('message_uuid') or message_uuid last_event_id = chunk.event_id yield chunk.encode() if chunk.event in ('done', 'error'): @@ -41,7 +44,7 @@ class SSEChatStreamService: except Exception: return - assistant = self._get_assistant_message() + assistant = self._get_assistant_message(message_uuid) if assistant and assistant.content: last_event_id += 1 yield SSEChunkService.done(last_event_id, assistant.content).encode() @@ -50,8 +53,11 @@ class SSEChatStreamService: yield SSEChunkService.error(last_event_id + 1, _('Stream timeout')).encode() - def _get_assistant_message(self) -> Message | None: - input_qs = Message.objects.filter(pk=self.store.message_uuid, is_deleted=False) + def _get_assistant_message(self, message_uuid) -> Message | None: + if not message_uuid: + return None + + input_qs = Message.objects.filter(pk=message_uuid, is_deleted=False) input_created_at = Subquery(input_qs.values('created_at')[:1]) return ( @@ -9,8 +9,7 @@ from tools.chats.domain import SSEChunk class SSEStoreService: - def __init__(self, message_uuid: UUID, user_uuid: UUID, chat_uuid: UUID) -> None: - self.message_uuid = message_uuid + def __init__(self, user_uuid: UUID, chat_uuid: UUID) -> None: self.user_uuid = user_uuid self.chat_uuid = chat_uuid self.redis_client = self._get_redis_client() @@ -20,7 +19,7 @@ class SSEStoreService: return caches['default']._cache.get_client(write=True) def _get_cache_key(self) -> str: - return f'sse:tokens:{self.user_uuid}:{self.chat_uuid}:{self.message_uuid}' + return f'sse:tokens:{self.user_uuid}:{self.chat_uuid}' def start(self) -> None: pipe = self.redis_client.pipeline() @@ -46,6 +45,16 @@ class SSEStoreService: def exists(self) -> bool: return bool(self.redis_client.exists(self._get_cache_key())) + def get_message_uuid(self) -> UUID | None: + raw = self.redis_client.lindex(self._get_cache_key(), 1) + if not raw: + return None + chunk = orjson.loads(raw) + if chunk.get('event') != 'start': + return None + uid = chunk.get('data', {}).get('message_uuid') + return UUID(str(uid)) if uid else None + def get_chunks(self, offset: int = 0) -> list[SSEChunk]: return [ SSEChunk(**orjson.loads(raw)) @@ -8,7 +8,7 @@ from tools.chats.services.sse_store import SSEStoreService @shared_task(soft_time_limit=570, time_limit=600) def event_stream_task(chat_uuid: str, message_uuid: str, user_uuid: str) -> None: - store = SSEStoreService(message_uuid, user_uuid, chat_uuid) + store = SSEStoreService(user_uuid=user_uuid, chat_uuid=chat_uuid) stream = None event_id = 0 @@ -52,6 +52,8 @@ MINIO_SECRET_KEY=testtest # CELERY CELERY_BROKER_URL=redis://cache-mdb:6379/0 CELERY_RESULT_BACKEND=redis://cache-mdb:6379/0 +CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT=630 +CELERY_WORKER_PREFETCH_MULTIPLIER=1 # EMAIL # For free hosts - https://www.wpoven.com/tools/free-smtp-server-for-testing @@ -110,7 +110,8 @@ services: - C_FORCE_ROOT=true - RELEASE - ENVIRONMENT - command: celery -A backend worker -l INFO --concurrency 3 + command: ["celery", "-A", "backend", "worker", "-l", "INFO", "--concurrency", "3"] + stop_grace_period: 10m deploy: replicas: 1 <<: [ *default-deploy ]