@@ -1,46 +1,95 @@ +import base64 +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 -def _resolve_chat_model(model_ref: str) -> NeuronModel: +def _parse_body(body: dict) -> MessageInSchema: + raw = body.get('input') + if raw is None: + raise HttpError(400, _('Missing required parameter: input')) + file, lines = None, [] + if isinstance(raw, str): + content = raw.strip() + elif not isinstance(raw, list): + raise HttpError(400, _('Invalid input payload')) + else: + for item in raw: + if not isinstance(item, dict): + continue + role = item.get('role', 'user').capitalize() + ic = item.get('content') + if isinstance(ic, str): + if t := ic.strip(): + lines.append(f'[{role}] {t}') + continue + for part in ic or []: + if not isinstance(part, dict): + continue + pt = part.get('type') + if pt in ('input_text', 'text') and (t := part.get('text')) and (t := str(t).strip()): + lines.append(f'[{role}] {t}') + elif pt in ('input_image', 'image_url') and file is None: + url = part.get('image_url') + url = url.get('url', '') if isinstance(url, dict) else str(url or '') + if 'base64' in url: + buf = BytesIO(base64.b64decode(url.split('base64', 1)[-1].lstrip(','))) + 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, + ) + content = '\n'.join(lines).strip() + if not content and not file: + raise HttpError(400, _('The request must not be empty')) + kwargs = {'content': content or '[User]'} + if file: + kwargs['file'] = file + if isinstance(md := body.get('metadata'), dict) and md: + info = dict(md) + for k in ('ttft', 'tbt'): + if isinstance(info.get(k), str): + try: + info[k] = float(info[k]) + except ValueError: + pass + kwargs['info'] = info + return MessageInSchema(**kwargs) + + +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 _require_idempotency_key(request) -> str: - if not (key := request.headers.get('Idempotency-Key', '').strip()): - raise HttpError(400, _('Idempotency-Key header is not provided')) - return key - - -def _api_user(request): +def _stream_context(request, response_id: str | None = None) -> tuple: 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, - ) + 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')) - return user + 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) @OpenAIErrorService.view @@ -51,46 +100,27 @@ def openai_responses_stream(request, body: dict): raise HttpError(400, _('Only streaming is supported')) if not (model_ref := body.get('model')): raise HttpError(400, _('You must provide a model parameter')) - - model = _resolve_chat_model(model_ref) - store = PublicSSEStoreService( - user_uuid=_api_user(request).pk, - idempotency_key=_require_idempotency_key(request), - ) - stream = OpenAIStreamService(store) - stream.cleanup_if_finished() - - response = public_stream_message(request, model.slug, OpenAIStreamService.user_message_body(body)) - meta = stream.init_meta(model_ref) - response.streaming_content = stream.event_stream(request, meta=meta) + 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 @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): if not stream: raise HttpError(400, _('Only streaming reconnect is supported')) - - user = _api_user(request) - store = PublicSSEStoreService( - user_uuid=user.pk, idempotency_key=response_id or _require_idempotency_key(request) - ) - if not store.exists(): + user, store, svc = _stream_context(request, response_id) + if not store.exists() and not svc.load_meta(): raise HttpError(404, _('Stream not found')) - - openai_stream = OpenAIStreamService(store) - meta = openai_stream.init_meta('', reconnect=True) + meta = svc.init_meta('', reconnect=True) if str(meta.get('user_uuid')) != str(user.pk): raise HttpError(404, _('Stream not found')) - return StreamingHttpResponse( - openai_stream.event_stream(request, reconnect=True, meta=meta), + svc.event_stream(request, meta, reconnect=True, starting_after=starting_after), content_type='text/event-stream', headers=OpenAIStreamService.SSE_HEADERS, ) @@ -101,14 +131,10 @@ 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) @@ -5,60 +5,18 @@ from ninja.errors import HttpError class OpenAIErrorService: - ERR_TYPES = { - 400: 'invalid_request_error', - 401: 'authentication_error', - 403: 'permission_error', - 404: 'invalid_request_error', - 409: 'invalid_request_error', - 501: 'api_error', - } - ERR_CODES = { - 'only streaming is supported': 'streaming_not_supported', - 'only streaming reconnect is supported': 'streaming_not_supported', - 'you must provide a model parameter': 'missing_required_parameter', - 'missing required parameter: input': 'missing_required_parameter', - 'model not found': 'model_not_found', - 'response not found': 'response_not_found', - 'stream not found': 'stream_not_found', - 'stream already in progress': 'stream_in_progress', - 'idempotency-key header is not provided': 'missing_idempotency_key', - 'no api key in authorization header': 'missing_api_key', - 'api key not found': 'invalid_api_key', - 'api key expired': 'expired_api_key', - 'api key limit exceeded': 'insufficient_quota', - 'api key is not available for this account type': 'account_not_allowed', - 'stream not supported for this model': 'streaming_not_supported', + TYPES = { + 400: 'invalid_request_error', 401: 'authentication_error', 403: 'permission_error', + 404: 'invalid_request_error', 409: 'invalid_request_error', 501: 'api_error', } @classmethod - def response( - cls, - status: int, - message, - *, - code: str | None = None, - param: str | None = None, - ) -> JsonResponse: - msg = str(message).strip() - if not code: - code = cls._infer_code(status, msg) + def response(cls, status: int, message: str) -> JsonResponse: return JsonResponse( - { - 'error': { - 'message': msg, - 'type': cls.ERR_TYPES.get(status, 'api_error'), - 'param': param, - 'code': code, - } - }, + {'error': {'message': str(message), 'type': cls.TYPES.get(status, 'api_error'), 'param': None, 'code': None}}, status=status, ) - @classmethod - def from_http_error(cls, exc: HttpError) -> JsonResponse: - return cls.response(exc.status_code, exc.message) - @classmethod def view(cls, fn): @wraps(fn) @@ -66,43 +24,6 @@ class OpenAIErrorService: try: return fn(*args, **kwargs) except HttpError as exc: - return cls.from_http_error(exc) + return cls.response(exc.status_code, exc.message) return wrapper - - @classmethod - def _infer_code(cls, status: int, message: str) -> str | None: - lower = message.lower() - if code := cls.ERR_CODES.get(lower): - return code - if status == 409: - return 'stream_in_progress' - if 'idempotency' in lower: - return 'missing_idempotency_key' - if status == 404: - return next( - ( - c - for k, c in ( - ('api key', 'invalid_api_key'), - ('model', 'model_not_found'), - ('response', 'response_not_found'), - ('stream', 'stream_not_found'), - ) - if k in lower - ), - None, - ) - if status == 403: - if 'limit' in lower: - return 'insufficient_quota' - if 'account' in lower: - return 'account_not_allowed' - if status == 501: - return 'streaming_not_supported' - if status == 401: - if 'expired' in lower: - return 'expired_api_key' - if 'api key' in lower: - return 'missing_api_key' - return None @@ -3,97 +3,39 @@ import time from typing import Iterator 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.schemas import MessageInSchema from tools.chats.services.sse_chat_stream import SSEChatStreamService from tools.chats.services.sse_store import PublicSSEStoreService class OpenAIStreamService: SSE_HEADERS = {'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'} - TERMINAL_EVENTS = frozenset({'done', 'error'}) - REASONING = {'effort': None, 'summary': None} - ITEM_IDX = {'output_index': 0, 'content_index': 0} - RESPONSE_DEFAULTS = { - 'parallel_tool_calls': True, - 'previous_response_id': None, - 'temperature': 1.0, - 'text': {'format': {'type': 'text'}}, - 'tool_choice': 'auto', - 'tools': [], - 'top_p': 1.0, - 'truncation': 'disabled', - 'user': None, - 'metadata': {}, - 'incomplete_details': None, - 'instructions': None, - 'max_output_tokens': None, - } + SETUP_N = 3 + IDX = {'output_index': 0, 'content_index': 0} def __init__(self, store: PublicSSEStoreService) -> None: self.store = store - @staticmethod - def parse_input(raw_input) -> str: - if raw_input is None: - raise HttpError(400, _('Missing required parameter: input')) - if isinstance(raw_input, str): - content = raw_input.strip() - elif not isinstance(raw_input, list): - raise HttpError(400, _('Invalid input payload')) - else: - lines = [] - for item in raw_input: - if not isinstance(item, dict): - continue - role = item.get('role', 'user').capitalize() - item_content = item.get('content') - if isinstance(item_content, str): - if text := item_content.strip(): - lines.append(f'[{role}] {text}') - continue - for part in item_content or []: - if ( - isinstance(part, dict) - and part.get('type') in ('input_text', 'text') - and (t := part.get('text')) - and (text := str(t).strip()) - ): - lines.append(f'[{role}] {text}') - content = '\n'.join(lines).strip() - if not content: - raise HttpError(400, _('The request must not be empty')) - return content - - @classmethod - def user_message_body(cls, body: dict) -> MessageInSchema: - kwargs = {'content': cls.parse_input(body.get('input'))} - if isinstance(body.get('metadata'), dict) and body['metadata']: - info = {} - for k, v in body['metadata'].items(): - if k in ('ttft', 'tbt') and isinstance(v, str): - try: - v = float(v) - except ValueError: - pass - info[k] = v - kwargs['info'] = info - return MessageInSchema(**kwargs) + 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 cleanup_if_finished(self) -> None: - if self.store.exists() and any( - orjson.loads(raw).get('event') in self.TERMINAL_EVENTS - for raw in self.store.redis_client.lrange(self.store._get_cache_key(), 0, -1) - ): - self._cleanup() + 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 = False) -> dict: - if meta := self._load_meta(): + def init_meta(self, model: str, *, reconnect: bool) -> dict: + if meta := self.load_meta(): return meta - if reconnect and not self.store.exists(): + if reconnect: raise HttpError(404, _('Stream not found')) meta = { 'rid': str(self.store.idempotency_key), @@ -101,104 +43,44 @@ class OpenAIStreamService: 'ts': int(time.time()), 'model': model, 'user_uuid': str(self.store.user_uuid), - 'last_seq': -1, - 'internal_offset': 0, - 'completed': False, - 'setup_emitted': False, } self._save_meta(meta) return meta def event_stream( self, - request: HttpRequest | None = None, + request: HttpRequest | None, + meta: dict, *, reconnect: bool = False, - meta: dict | None = None, - model: str = '', + starting_after: int = 0, ) -> Iterator[str]: - meta = meta or self.init_meta(model, reconnect=reconnect) - replayed: list[str] = [] - if reconnect: - if meta.get('completed'): - try: - yield from self._replay_events() - finally: - self._cleanup() - return - replayed = self._replay_events() - yield from replayed - if (self._load_meta() or {}).get('completed'): - self._cleanup() - return - if not reconnect or not replayed: - yield from self._emit_setup(meta) - try: - yield from self._map_internal_stream(meta, request) - except Exception as exc: - yield from self._finish_error(meta, exc) - self._cleanup() - - def _access(self) -> tuple[int, dict, dict]: - key = self.store._get_cache_key() - raws = self.store.redis_client.lrange(key, 0, -1) - if not raws: - chunk = {'event_id': 0, 'event': 'pending', 'data': {}} - self.store.redis_client.rpush(key, orjson.dumps(chunk)) - index = 0 - else: - index, chunk = 0, None - for i, raw in enumerate(raws): - chunk = orjson.loads(raw) - if chunk.get('event') == 'pending' or (chunk.get('data') or {}).get('openai') is not None: - index = i - break - else: - chunk = {'event_id': 0, 'event': 'pending', 'data': {}} - index = len(raws) - self.store.redis_client.rpush(key, orjson.dumps(chunk)) - bucket = chunk.setdefault('data', {}).setdefault('openai', {'meta': None, 'events': []}) - return index, chunk, bucket - - def _commit(self, index: int, chunk: dict) -> None: - self.store.redis_client.lset(self.store._get_cache_key(), index, orjson.dumps(chunk)) + 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) - def _load_meta(self) -> dict | None: - if not self.store.exists(): - return None - return self._access()[2].get('meta') - - def _save_meta(self, meta: dict) -> None: + def _finalize_session(self) -> None: + self._del_meta() if self.store.exists(): - index, chunk, bucket = self._access() - bucket['meta'] = meta - self._commit(index, chunk) + self.store.cleanup() - def _replay_events(self) -> list[str]: - if not self.store.exists(): - return [] - events = self._access()[2].get('events') or [] - return [e if isinstance(e, str) else e['sse'] for e in events] + def _meta_key(self) -> str: + return f'sse:openai:meta:{self.store.user_uuid}:{self.store.idempotency_key}' - def _cleanup(self) -> None: - if self.store.exists(): - self.store.cleanup() + def _save_meta(self, meta: dict) -> None: + self.store.redis_client.set(self._meta_key(), orjson.dumps(meta), ex=settings.SSE_STREAM_TTL) - @staticmethod - def _omit_none(value): - if isinstance(value, dict): - return {k: OpenAIStreamService._omit_none(v) for k, v in value.items() if v is not None} - if isinstance(value, list): - return [OpenAIStreamService._omit_none(v) for v in value] - return value + def _del_meta(self) -> None: + self.store.redis_client.delete(self._meta_key()) @classmethod def _sse(cls, event_type: str, seq: int, **fields) -> str: - payload = cls._omit_none({'type': event_type, 'sequence_number': seq, **fields}) + payload = {'type': event_type, 'sequence_number': seq, **fields} return f'event: {event_type}\ndata: {orjson.dumps(payload).decode()}\n\n' @staticmethod - def _parse_sse(chunk: str) -> tuple[str, dict, int | None]: + def _parse_chunk(chunk: str) -> tuple[str, dict, int | None]: event, data, event_id = '', {}, None for line in chunk.split('\n'): if line.startswith('id:'): @@ -209,112 +91,90 @@ class OpenAIStreamService: data = orjson.loads(line[5:].strip()) return event, data, event_id - def _response_obj(self, meta: dict, status: str, **fields): + @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): return { 'id': meta['rid'], 'object': 'response', 'created_at': meta['ts'], 'model': meta['model'], 'status': status, - 'error': None, - **self.RESPONSE_DEFAULTS, - **fields, + 'output': output, + 'store': True, + 'text': {'format': {'type': 'text'}}, } - def _emit(self, meta: dict, event_type: str, **fields) -> str: - meta['last_seq'] += 1 - sse = self._sse(event_type, meta['last_seq'], **fields) - index, chunk, bucket = self._access() - bucket['events'].append(sse) - self._commit(index, chunk) - return sse - - def _mark_done(self, meta: dict) -> None: - meta['completed'] = True - self._save_meta(meta) - - def _emit_setup(self, meta: dict) -> Iterator[str]: - if meta.get('setup_emitted'): - return - progress = self._response_obj(meta, 'in_progress', output=[], reasoning=self.REASONING, store=True, usage=None) - yield self._emit(meta, 'response.created', response=progress) - yield self._emit(meta, 'response.in_progress', response=progress) - yield self._emit( - meta, - 'response.output_item.added', - item={'id': meta['mid'], 'type': 'message', 'status': 'in_progress', 'role': 'assistant', 'content': []}, - **self.ITEM_IDX, - ) - yield self._emit( - meta, - 'response.content_part.added', - item_id=meta['mid'], - part={'type': 'output_text', 'text': '', 'annotations': []}, - **self.ITEM_IDX, - ) - meta['setup_emitted'] = True - self._save_meta(meta) - - def _finish_success(self, meta: dict, text: str) -> Iterator[str]: + 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, + }), + )): + if seq > min_seq: + yield self._sse(event_type, seq, **fields) + + def _closing(self, meta: dict, text: str, first_seq: int, min_seq: int) -> 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.ITEM_IDX} - yield self._emit(meta, 'response.output_text.done', text=text, **ctx) - yield self._emit(meta, 'response.content_part.done', part=part, **ctx) - yield self._emit(meta, 'response.output_item.done', item=item, output_index=0) - yield self._emit( - meta, - 'response.completed', - response=self._response_obj( - meta, - 'completed', - output=[item], - reasoning=self.REASONING, - store=True, - usage={ - 'input_tokens': 0, - 'output_tokens': 0, - 'output_tokens_details': {'reasoning_tokens': 0}, - 'total_tokens': 0, - }, - ), - ) - self._mark_done(meta) - - def _finish_error(self, meta: dict, message) -> Iterator[str]: - yield self._emit( - meta, - 'response.failed', - response=self._response_obj( - meta, 'failed', store=False, error={'code': 'server_error', 'message': str(message)} - ), - ) - self._mark_done(meta) + 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])}), + ): + if seq > min_seq: + yield self._sse(event_type, seq, **fields) + seq += 1 - def _map_internal_stream(self, meta: dict, request: HttpRequest | None) -> Iterator[str]: - text_parts = [] - ctx = {'item_id': meta['mid'], **self.ITEM_IDX} - try: - for chunk in SSEChatStreamService(self.store).event_stream( - request=request, offset=int(meta.get('internal_offset', 0)) - ): - if chunk == SSEChatStreamService.HEARTBEAT: - continue - event, data, event_id = self._parse_sse(chunk) - if event == 'token' and (token := data.get('content', '')): - text_parts.append(token) - yield self._emit(meta, 'response.output_text.delta', delta=token, **ctx) - elif event == 'done': - yield from self._finish_success(meta, data.get('content') or ''.join(text_parts)) - return - elif event == 'error': - yield from self._finish_error( - meta, data.get('detail', 'The model failed to generate a response.') + 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', ''))}, + }, ) - return - if event_id is not None: - meta['internal_offset'] = int(event_id) - self._save_meta(meta) - finally: - if meta.get('completed'): - self._cleanup() + self._finalize_session() + return