@@ -69,10 +69,11 @@ class SSEStoreService: class PublicSSEStoreService(SSEStoreService): - def __init__(self, idempotency_key: UUID, user_uuid: UUID): + def __init__(self, message_uuid: UUID, user_uuid: UUID): self.user_uuid = user_uuid - self.idempotency_key = idempotency_key + self.message_uuid = message_uuid self.redis_client = self._get_redis_client() def _get_cache_key(self) -> str: - return f'sse:tokens:{self.user_uuid}:{self.idempotency_key}' \ No newline at end of file + return f'sse:tokens:{self.user_uuid}:{self.message_uuid}:public' + \ No newline at end of file @@ -61,11 +61,10 @@ def public_event_stream_task( message_uuid: str, user_uuid: str, model_slug: str, - idempotency_key: str, api_key_uuid: str, debit_api_key_limit: bool, ): - store = PublicSSEStoreService(user_uuid=user_uuid, idempotency_key=idempotency_key) + store = PublicSSEStoreService(user_uuid=user_uuid, message_uuid=message_uuid) message = Message.objects.get(pk=message_uuid) api_store = APIStore.objects.select_related( @@ -3,14 +3,12 @@ from io import BytesIO import filetype from django.db.models import Q -from django.http import StreamingHttpResponse from django.core.files.uploadedfile import InMemoryUploadedFile from django.utils.translation import gettext as _ from ninja.errors import HttpError from ml_model.models import NeuronModel from tools.chats.schemas import MessageInSchema -from tools.chats.services.sse_store import PublicSSEStoreService from tools.public_api.services.openai_errors import OpenAIErrorService from tools.public_api.services.openai_stream import OpenAIStreamService @@ -48,8 +46,12 @@ def _parse_body(body: dict) -> MessageInSchema: kind = filetype.guess(buf.read(20)) buf.seek(0) file = InMemoryUploadedFile( - buf, 'file', f'api-file.{kind.extension if kind else "bin"}', - kind.mime if kind else 'application/octet-stream', buf.getbuffer().nbytes, None, + buf, + 'file', + f'api-file.{kind.extension if kind else "bin"}', + kind.mime if kind else 'application/octet-stream', + buf.getbuffer().nbytes, + None, ) content = '\n'.join(lines).strip() if not content and not file: @@ -71,25 +73,41 @@ def _parse_body(body: dict) -> MessageInSchema: def _resolve_model(model_ref: str) -> NeuronModel: try: - return NeuronModel.objects.filter( - Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), - category__slug='chat-bots', - ).distinct().get() + return ( + NeuronModel.objects.filter( + Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), + category__slug='chat-bots', + ) + .distinct() + .get() + ) except NeuronModel.DoesNotExist: raise HttpError(404, _('Model not found')) -def _stream_context(request, response_id: str | None = None) -> tuple: +def _public_user_uuid(request): from tools.public_api.routes.v1 import _get_api_key - api_key = _get_api_key(request, select_related=['user__host_account', 'user__business_account'], check_usage_limit=False) - user = api_key.user - if user.account_type not in ('business_host', 'regular', 'business_admin'): - raise HttpError(403, _('API key is not available for this account type')) - if not (key := (response_id or request.headers.get('Idempotency-Key', '')).strip()): - raise HttpError(400, _('Idempotency-Key header is not provided')) - store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=key) - return user, store, OpenAIStreamService(store) + return _get_api_key(request, check_usage_limit=False).user.pk + + +def _to_openai( + response, + model_ref: str, + *, + request, + message_uuid: str | None = None, + starting_after: int = 0, +): + # Django 6: .streaming_content yields bytes; wrap the raw str iterator instead. + response.streaming_content = OpenAIStreamService.translate( + response._iterator, + model_ref, + message_uuid=message_uuid, + starting_after=starting_after, + user_uuid=_public_user_uuid(request), + ) + return response @OpenAIErrorService.view @@ -101,28 +119,30 @@ def openai_responses_stream(request, body: dict): if not (model_ref := body.get('model')): raise HttpError(400, _('You must provide a model parameter')) model = _resolve_model(model_ref) - _, store, svc = _stream_context(request) - svc.cleanup_finished() - response = public_stream_message(request, model.slug, _parse_body(body)) - meta = svc.init_meta(model_ref, reconnect=False) - response.streaming_content = svc.event_stream(request, meta) - return response + return _to_openai(public_stream_message(request, model.slug, _parse_body(body)), model_ref, request=request) @OpenAIErrorService.view -def openai_responses_stream_reconnect(request, response_id: str | None = None, *, stream: bool = True, starting_after: int = 0): +def openai_responses_stream_reconnect( + request, + response_id: str | None = None, + *, + stream: bool = True, + starting_after: int = 0, +): + from tools.public_api.routes.v1 import public_stream_message_reconnect + if not stream: raise HttpError(400, _('Only streaming reconnect is supported')) - user, store, svc = _stream_context(request, response_id) - if not store.exists() and not svc.load_meta(): - raise HttpError(404, _('Stream not found')) - meta = svc.init_meta('', reconnect=True) - if str(meta.get('user_uuid')) != str(user.pk): - raise HttpError(404, _('Stream not found')) - return StreamingHttpResponse( - svc.event_stream(request, meta, reconnect=True, starting_after=starting_after), - content_type='text/event-stream', - headers=OpenAIStreamService.SSE_HEADERS, + if not (message_uuid := (response_id or request.GET.get('response_id', '')).strip()): + raise HttpError(400, _('message_uuid is not provided')) + response = public_stream_message_reconnect( + request, + message_uuid, + offset=OpenAIStreamService.public_offset(starting_after), + ) + return _to_openai( + response, '', message_uuid=message_uuid, starting_after=starting_after, request=request ) @@ -131,10 +151,14 @@ from tools.public_api.routes import v1 as public_v1_routes def _register_reconnect_get(path: str, *, with_response_id: bool): if with_response_id: + @public_v1_routes.router.get(path, tags=['openai/responses']) def handler(request, response_id: str, stream: bool = True, starting_after: int = 0): - return openai_responses_stream_reconnect(request, response_id, stream=stream, starting_after=starting_after) + return openai_responses_stream_reconnect( + request, response_id, stream=stream, starting_after=starting_after + ) else: + @public_v1_routes.router.get(path, tags=['openai/responses']) def handler(request, stream: bool = True, starting_after: int = 0): return openai_responses_stream_reconnect(request, stream=stream, starting_after=starting_after) @@ -50,8 +50,8 @@ def _get_api_key( return api_key -@router.get('text/{model_slug}/stream', tags=['public/text']) -def public_stream_message_reconnect(request, model_slug: str, offset: int = 0): +@router.get('text/{message_uuid}/stream/reconnect', tags=['public/text']) +def public_stream_message_reconnect(request, message_uuid: str, offset: int = 0): api_key = _get_api_key( request, select_related=['user__host_account', 'user__business_account'], @@ -61,11 +61,7 @@ def public_stream_message_reconnect(request, model_slug: str, offset: int = 0): if user.account_type not in ('business_host', 'regular', 'business_admin'): raise HttpError(403, _('API key is not available for this account type')) - idempotency_key = request.headers.get('Idempotency-Key', '') - if not idempotency_key: - raise HttpError(400, _('Idempotency-Key header is not provided')) - - store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=idempotency_key) + store = PublicSSEStoreService(user_uuid=user.pk, message_uuid=message_uuid) if not store.exists(): raise HttpError(404, _('Stream not found')) @@ -100,10 +96,6 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): raise HttpError(403, _('API key is not available for this account type')) balance = user.balance - idempotency_key = request.headers.get('Idempotency-Key', '') - if not idempotency_key: - raise HttpError(400, _('Idempotency-Key header is not provided')) - api_store, created = APIStore.objects.get_or_create(user=user) selector = NeuronModelSelector(user) @@ -116,14 +108,11 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): if not body.content: raise HttpError(400, _('The request must not be empty')) - store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=idempotency_key) - 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=api_store, from_model=False, from_public_api=True, **data ) + store = PublicSSEStoreService(user_uuid=user.pk, message_uuid=input_message.pk) store.start() public_event_stream_task.delay( @@ -131,7 +120,6 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): message_uuid=str(input_message.pk), user_uuid=str(user.pk), model_slug=model_slug, - idempotency_key=idempotency_key, api_key_uuid=str(api_key.pk), debit_api_key_limit=api_key.token_limit is not None, ) @@ -1,12 +1,8 @@ -import secrets import time from typing import Iterator +from uuid import UUID import orjson -from django.conf import settings -from django.http import HttpRequest -from django.utils.translation import gettext as _ -from ninja.errors import HttpError from tools.chats.services.sse_chat_stream import SSEChatStreamService from tools.chats.services.sse_store import PublicSSEStoreService @@ -15,64 +11,82 @@ from tools.chats.services.sse_store import PublicSSEStoreService class OpenAIStreamService: SSE_HEADERS = {'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'} SETUP_N = 3 - IDX = {'output_index': 0, 'content_index': 0} - - def __init__(self, store: PublicSSEStoreService) -> None: - self.store = store - - def cleanup_finished(self) -> None: - if not self.store.exists(): - self._del_meta() - return - key = self.store._get_cache_key() - if any(orjson.loads(r).get('event') in ('done', 'error') for r in self.store.redis_client.lrange(key, 0, -1)): - self.store.cleanup() - self._del_meta() - - def load_meta(self) -> dict | None: - return orjson.loads(r) if (r := self.store.redis_client.get(self._meta_key())) else None - - def init_meta(self, model: str, *, reconnect: bool) -> dict: - if meta := self.load_meta(): - return meta - if reconnect: - raise HttpError(404, _('Stream not found')) - meta = { - 'rid': str(self.store.idempotency_key), - 'mid': f'msg_{secrets.token_hex(24)}', - 'ts': int(time.time()), - 'model': model, - 'user_uuid': str(self.store.user_uuid), - } - self._save_meta(meta) - return meta + _IDX = {'output_index': 0, 'content_index': 0} + + @classmethod + def public_offset(cls, starting_after: int) -> int: + return 0 if starting_after < cls.SETUP_N - 1 else starting_after - cls.SETUP_N + 2 - def event_stream( - self, - request: HttpRequest | None, - meta: dict, + @classmethod + def translate( + cls, + public_stream, + model_ref: str, *, - reconnect: bool = False, + message_uuid: str | None = None, starting_after: int = 0, + user_uuid: UUID | None = None, ) -> Iterator[str]: - min_seq = -1 if starting_after == 0 else starting_after - offset = 0 if starting_after < self.SETUP_N - 1 else starting_after - self.SETUP_N + 2 - emit_setup = not reconnect or starting_after < self.SETUP_N - 1 - yield from self._translate(request, meta, offset=offset, min_seq=min_seq, emit_setup=emit_setup) + min_seq = starting_after or -1 + emit_setup = starting_after < cls.SETUP_N - 1 + meta = cls._meta(message_uuid, model_ref) if message_uuid else None + setup_sent = emit_setup and bool(meta) + tokens: list[str] = [] + last_seq = min_seq - def _finalize_session(self) -> None: - self._del_meta() - if self.store.exists(): - self.store.cleanup() + if setup_sent: + yield from cls._setup(meta, min_seq) - def _meta_key(self) -> str: - return f'sse:openai:meta:{self.store.user_uuid}:{self.store.idempotency_key}' + for chunk in public_stream: + if chunk == SSEChatStreamService.HEARTBEAT: + continue + event, data, eid = cls._parse(chunk) - def _save_meta(self, meta: dict) -> None: - self.store.redis_client.set(self._meta_key(), orjson.dumps(meta), ex=settings.SSE_STREAM_TTL) + if event == 'start': + meta = meta or cls._meta(data['message_uuid'], model_ref) + if emit_setup and not setup_sent: + yield from cls._setup(meta, min_seq) + setup_sent = True + continue + if event == 'pending' or not meta: + continue - def _del_meta(self) -> None: - self.store.redis_client.delete(self._meta_key()) + ctx = {'item_id': meta['mid'], **cls._IDX} + if event == 'token' and eid and (token := data.get('content', '')): + tokens.append(token) + if (seq := cls.SETUP_N + eid - 2) > min_seq: + yield cls._sse('response.output_text.delta', seq, delta=token, **ctx) + last_seq = seq + elif event == 'done': + text = data.get('content') or ''.join(tokens) + first = cls.SETUP_N + eid - 2 if eid else cls.SETUP_N + yield from cls._close(meta, text, first, min_seq, ctx) + cls._cleanup_public_stream(user_uuid, meta['rid']) + return + elif event == 'error': + if (seq := last_seq + 1) > min_seq: + yield cls._sse( + 'response.failed', + seq, + response={ + 'id': meta['rid'], + 'object': 'response', + 'status': 'failed', + 'error': {'code': 'server_error', 'message': str(data.get('detail', ''))}, + }, + ) + cls._cleanup_public_stream(user_uuid, meta['rid']) + return + + @classmethod + def _cleanup_public_stream(cls, user_uuid: UUID | None, message_uuid: str | None) -> None: + if user_uuid and message_uuid: + PublicSSEStoreService(user_uuid=user_uuid, message_uuid=message_uuid).cleanup() + + @classmethod + def _meta(cls, message_uuid: str, model: str) -> dict: + uid = str(message_uuid) + return {'rid': uid, 'mid': f'msg_{uid.replace("-", "")}', 'ts': int(time.time()), 'model': model} @classmethod def _sse(cls, event_type: str, seq: int, **fields) -> str: @@ -80,22 +94,19 @@ class OpenAIStreamService: return f'event: {event_type}\ndata: {orjson.dumps(payload).decode()}\n\n' @staticmethod - def _parse_chunk(chunk: str) -> tuple[str, dict, int | None]: - event, data, event_id = '', {}, None + def _parse(chunk: str) -> tuple[str, dict, int | None]: + event, data, eid = '', {}, None for line in chunk.split('\n'): if line.startswith('id:'): - event_id = int(line[3:].strip()) + eid = int(line[3:].strip()) elif line.startswith('event:'): event = line[6:].strip() elif line.startswith('data:'): data = orjson.loads(line[5:].strip()) - return event, data, event_id + return event, data, eid @classmethod - def _token_seq(cls, event_id: int) -> int: - return cls.SETUP_N + (event_id - 2) - - def _response(self, meta: dict, status: str, output: list): + def _resp(cls, meta: dict, status: str, output: list) -> dict: return { 'id': meta['rid'], 'object': 'response', @@ -107,74 +118,47 @@ class OpenAIStreamService: 'text': {'format': {'type': 'text'}}, } - def _setup(self, meta: dict, min_seq: int) -> Iterator[str]: - item = {'id': meta['mid'], 'type': 'message', 'status': 'in_progress', 'role': 'assistant', 'content': []} - created = self._response(meta, 'in_progress', []) - for seq, (event_type, fields) in enumerate(( - ('response.created', {'response': created}), - ('response.output_item.added', {'item': item, **self.IDX}), - ('response.content_part.added', { - 'item_id': meta['mid'], - 'part': {'type': 'output_text', 'text': '', 'annotations': []}, - **self.IDX, - }), - )): + @classmethod + def _setup(cls, meta: dict, min_seq: int) -> Iterator[str]: + item = { + 'id': meta['mid'], + 'type': 'message', + 'status': 'in_progress', + 'role': 'assistant', + 'content': [], + } + for seq, (etype, fields) in enumerate( + ( + ('response.created', {'response': cls._resp(meta, 'in_progress', [])}), + ('response.output_item.added', {'item': item, **cls._IDX}), + ( + 'response.content_part.added', + { + 'item_id': meta['mid'], + 'part': {'type': 'output_text', 'text': '', 'annotations': []}, + **cls._IDX, + }, + ), + ) + ): if seq > min_seq: - yield self._sse(event_type, seq, **fields) + yield cls._sse(etype, seq, **fields) - def _closing(self, meta: dict, text: str, first_seq: int, min_seq: int) -> Iterator[str]: + @classmethod + def _close(cls, meta: dict, text: str, first_seq: int, min_seq: int, ctx: dict) -> Iterator[str]: part = {'type': 'output_text', 'text': text, 'annotations': []} - item = {'id': meta['mid'], 'type': 'message', 'status': 'completed', 'role': 'assistant', 'content': [part]} - ctx = {'item_id': meta['mid'], **self.IDX} - seq = first_seq - for event_type, fields in ( - ('response.output_text.done', {'text': text, **ctx}), - ('response.completed', {'response': self._response(meta, 'completed', [item])}), + item = { + 'id': meta['mid'], + 'type': 'message', + 'status': 'completed', + 'role': 'assistant', + 'content': [part], + } + for i, (etype, fields) in enumerate( + ( + ('response.output_text.done', {'text': text, **ctx}), + ('response.completed', {'response': cls._resp(meta, 'completed', [item])}), + ) ): - if seq > min_seq: - yield self._sse(event_type, seq, **fields) - seq += 1 - - def _translate( - self, - request: HttpRequest | None, - meta: dict, - *, - offset: int, - min_seq: int, - emit_setup: bool, - ) -> Iterator[str]: - if emit_setup: - yield from self._setup(meta, min_seq) - tokens, ctx, last_seq = [], {'item_id': meta['mid'], **self.IDX}, min_seq - for chunk in SSEChatStreamService(self.store).event_stream(request=request, offset=offset): - if chunk == SSEChatStreamService.HEARTBEAT: - continue - event, data, event_id = self._parse_chunk(chunk) - if event == 'token' and event_id and (token := data.get('content', '')): - tokens.append(token) - seq = self._token_seq(event_id) - if seq > min_seq: - yield self._sse('response.output_text.delta', seq, delta=token, **ctx) - last_seq = seq - elif event == 'done': - text = data.get('content') or ''.join(tokens) - close_seq = self._token_seq(event_id - 1) + 1 if event_id else self.SETUP_N - yield from self._closing(meta, text, close_seq, min_seq) - self._finalize_session() - return - elif event == 'error': - err_seq = last_seq + 1 - if err_seq > min_seq: - yield self._sse( - 'response.failed', - err_seq, - response={ - 'id': meta['rid'], - 'object': 'response', - 'status': 'failed', - 'error': {'code': 'server_error', 'message': str(data.get('detail', ''))}, - }, - ) - self._finalize_session() - return + if (seq := first_seq + i) > min_seq: + yield cls._sse(etype, seq, **fields)