@@ -325,6 +325,8 @@ UPSCALE_MULTIPLIER_HOST = env.str('UPSCALE_MULTIPLIER_HOST', 'packet:8080') # FILES ImageFile.LOAD_TRUNCATED_IMAGES = True +DATA_UPLOAD_MAX_MEMORY_SIZE = env.int('DATA_UPLOAD_MAX_MEMORY_SIZE', 5) << 20 + # Payments YOOKASSA_ACCOUNT_ID = env.str('YOOKASSA_ACCOUNT_ID', default='defaultapikey') YOOKASSA_SECRET_KEY = env.str('YOOKASSA_SECRET_KEY', default='defaultapikey') @@ -488,3 +490,6 @@ UNLEASH_WEBHOOK_SECRET_KEY = env.str('UNLEASH_WEBHOOK_SECRET_KEY', 'defaultsecre # RECURRING SETTINGS MAX_RECURRING_ATTEMPTS = env.int('MAX_RECURRING_ATTEMPTS', 1) + +# SSE STREAMING +FF__STREAMING_ENABLED = env.bool('FF__STREAMING_ENABLED', False) @@ -4,7 +4,7 @@ from typing import Literal from django.conf import settings from django.conf.urls.static import static from django.contrib import admin -from django.core.exceptions import ObjectDoesNotExist +from django.core.exceptions import ObjectDoesNotExist, RequestDataTooBig from django.urls import include, path from django.utils.translation import gettext_lazy as _ from drf_spectacular.views import SpectacularAPIView, SpectacularSwaggerView @@ -79,6 +79,14 @@ def object_does_not_exists_error_handler(request, exc: ObjectDoesNotExist): return api.create_response(request, {'message': _('Requested object does not exists')}, status=404) +@compatibility_api.exception_handler(RequestDataTooBig) +@api.exception_handler(RequestDataTooBig) +def request_data_too_big_error_handler(request, exc: RequestDataTooBig): + return api.create_response( + request, {'detail': _('Data size limit exceeded. Please reduce the size')}, status=413 + ) + + @api.exception_handler(InvalidToken) def invalid_token_error_handler(request, exc: InvalidToken): return api.create_response(request, {'message': _('Token is invalid')}, status=401) @@ -1178,6 +1178,10 @@ msgstr "Списания" msgid "Amount" msgstr "Количество" +#: lib/middleware.py:40 +msgid "Data size limit exceeded. Please reduce the size" +msgstr "Превышен лимит размера данных. Уменьшите размер" + #: payments/models/payment.py:41 msgid "Plan" msgstr "План" @@ -81,3 +81,6 @@ from ml_model.services.vicuna import Vicuna from ml_model.services.wan import Wan from ml_model.services.wan_lite import Wan_Lite from ml_model.services.whisper import Whisper +from ml_model.services.image_test_model import Image_Test_Model +from ml_model.services.video_test_model import Video_Test_Model +from ml_model.services.audio_test_model import Audio_Test_Model \ No newline at end of file @@ -0,0 +1,97 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any +import filetype + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.exceptions import CorruptedFileError, InvalidParameterError +from ml_model.services.base import SimpleService +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + +import random + +class Audio_Test_Model(SimpleService): + + TOKENS_COST = Decimal('3') + + PLACEHOLDER_URL=[ + 'https://www.myinstants.com/media/sounds/saliut-eblany-batia-doma-billy-butcher-i-the-boys.mp3', + 'https://www.myinstants.com/media/sounds/zdravstvuite-nichtozhnye-nishchie-smertnye.mp3', + 'https://www.myinstants.com/media/sounds/okh-zria-ia-tuda-polez.mp3'] # Позже убрать + + def calculate_price(self, num_audios: int = 1) -> Decimal: + return self.TOKENS_COST * num_audios + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + num_audios = info.get('num_audios', 1) + + return cls.TOKENS_COST * num_audios + + def save_results( + self, + content: str, + t: timedelta, + audios: list[bytes], + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + for audio in audios: + messages.append( + Message( + content_object=self.store, + elapsed_time=t, + content=content, + file=File(BytesIO(audio), '.mp3'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + cau = input_message.info.get('cau') or self.PLACEHOLDER_URL[random.randint(0, 2)] # Позже убрать + num_audios = input_message.info.get('num_audios', 1) + + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < (cost := self.calculate_price(num_audios)): + raise InsufficientBalance(balance, cost) + + start_time = time.time() + + audio_bytes = self._fetch_audio(cau) + + audios = [audio_bytes] * num_audios + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, num_audios) + + msgs = self.save_results(input_message.content, process_time, audios, save) + return msgs + + + def _fetch_audio(self, url: str): + try: + response = requests.get( + url, + timeout=600 + ) + response.raise_for_status() + except requests.RequestException as exc: + raise InvalidParameterError(f'Invalid audio URL: {exc}') + + kind = filetype.guess(response.content[:120]) + + if not kind: + raise CorruptedFileError + + if not kind.mime.startswith('audio/'): + raise InvalidParameterError('Audio format not supported') + + return response.content \ No newline at end of file @@ -0,0 +1,90 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any +import filetype + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.exceptions import CorruptedFileError, InvalidParameterError +from ml_model.services.base import SimpleService +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + + +class Image_Test_Model(SimpleService): + + TOKENS_COST = Decimal('3') + PLACEHOLDER_URL='https://i.pinimg.com/736x/8b/e0/61/8be06158da3986fb4c47497b5660bb29.jpg' + + def calculate_price(self, num_images: int = 1 ) -> Decimal: + return self.TOKENS_COST * num_images + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + num_images = info.get('num_images', 1) + + return cls.TOKENS_COST * num_images + + def save_results( + self, + content: str, + t: timedelta, + images: list[bytes], + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + for image in images: + messages.append( + Message( + content_object=self.store, + elapsed_time=t, + content=content, + file=File(BytesIO(image), '.png'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + num_images = input_message.info.get('num_images', 1) + + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < (cost := self.TOKENS_COST * num_images): + raise InsufficientBalance(balance, cost) + + start_time = time.time() + ciu = input_message.info.get('ciu') or self.PLACEHOLDER_URL + + images = [self._fetch_image(ciu)] * num_images + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, num_images) + + msgs = self.save_results(input_message.content, process_time, images, save) + return msgs + + + + def _fetch_image(self, url: str): + try: + response = requests.get(url, timeout=10) + response.raise_for_status() + except requests.RequestException: + raise InvalidParameterError('Invalid image URL') + + kind = filetype.guess(response.content[:20]) + + if not kind: + raise CorruptedFileError + + if not kind.mime.startswith('image/'): + raise InvalidParameterError('Image format not supported') + + return response.content + \ No newline at end of file @@ -0,0 +1,140 @@ +import math +import subprocess +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any, Literal +import filetype +import json + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.exceptions import CorruptedFileError, InvalidParameterError +from ml_model.services.base import SimpleService +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + + +class Video_Test_Model(SimpleService): + + RATES = { + 'per-unit': Decimal('15'), + 'per-second': Decimal('0.5'), + } + + + PLACEHOLDER_URL='https://imgur.com/QPLhtj1.mp4' + + def calculate_price(self, strategy: Literal['per-unit','per-second'], duration: int = 1, num_videos: int = 1) -> Decimal: + rate = self.RATES[strategy] + return (rate * duration * num_videos).quantize(Decimal('0.1'), rounding='ROUND_UP') + + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + num_videos = info.get('num_videos', 1) + cps = info.get('cps', 'per-unit') + + if cps == 'per-second': + return None + + return cls.RATES['per-unit'] * num_videos + + + def save_results( + self, + content: str, + t: timedelta, + videos: list[bytes], + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + for video in videos: + messages.append( + Message( + content_object=self.store, + elapsed_time=t, + content=content, + file=File(BytesIO(video), '.mp4'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + cps = input_message.info.get('cps', 'per-unit') + cvu = input_message.info.get('cvu') or self.PLACEHOLDER_URL + num_videos = input_message.info.get('num_videos', 1) + + start_time = time.time() + + video_bytes = self._fetch_video(cvu) + + if cps == 'per-second': + duration = self._get_duration(video_bytes) + else: + duration = 1 + + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < (cost := self.calculate_price(cps, duration, num_videos)): + raise InsufficientBalance(balance, cost) + + + videos = [video_bytes] * num_videos + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, cps, duration, num_videos) + + msgs = self.save_results(input_message.content, process_time, videos, save) + return msgs + + + def _fetch_video(self, url: str): + headers = { + "User-Agent": ( + "Mozilla/5.0 (Windows NT 10.0; Win64; x64) " + "AppleWebKit/537.36 (KHTML, like Gecko) " + "Chrome/137.0 Safari/537.36" + ) + } + try: + response = requests.get( + url, + headers=headers, + timeout=600 + ) + response.raise_for_status() + except requests.RequestException as exc: + raise InvalidParameterError(f'Invalid video URL: {exc}') + + kind = filetype.guess(response.content[:120]) + + if not kind: + raise CorruptedFileError + + if not kind.mime.startswith('video/'): + raise InvalidParameterError('Video format not supported') + + return response.content + + + def _get_duration(self, video_bytes: bytes) -> int: + result = subprocess.run( + [ + "ffprobe", + "-v", "quiet", + "-print_format", "json", + "-show_format", + "-", + ], + input=video_bytes, + capture_output=True, + ) + + data = json.loads(result.stdout) + return math.ceil(float(data["format"]["duration"])) \ No newline at end of file @@ -1,6 +1,7 @@ import importlib from uuid import uuid4 +from django.conf import settings from django.contrib.contenttypes.fields import GenericForeignKey from django.contrib.contenttypes.models import ContentType from django.contrib.postgres.fields import ArrayField @@ -147,7 +148,7 @@ class NeuronModel(BaseModel, OrderedModel): @property def streaming(self) -> bool: - return hasattr(self.service, 'make_stream') + return settings.FF__STREAMING_ENABLED and hasattr(self.service, 'make_stream') def __str__(self): return self.title @@ -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) @@ -111,4 +111,9 @@ BUILDKIT_PROGRESS=plain # SERVER DJANGO_RUNSERVER_HIDE_WARNING=true -PYTHONWARNINGS=ignore::UserWarning:polymorphic # temporarily \ No newline at end of file +PYTHONWARNINGS=ignore::UserWarning:polymorphic # temporarily + +# SSE STREAMING +FF__STREAMING_ENABLED=True + +DATA_UPLOAD_MAX_MEMORY_SIZE=5 # MB \ No newline at end of file