@@ -2,7 +2,15 @@ from authentication.exceptions.business_host_exceptions.base_already import ( BaseAlready, ) -__all__ = ('BaseAlready',) +from django.utils.translation import gettext_lazy as _ + +__all__ = ( + 'BaseAlready', + 'InvalidPassword', + 'InvalidUsername', + 'InvalidToken', + 'APIException' +) class InvalidPassword(Exception): ... @@ -14,3 +22,12 @@ class InvalidUsername(Exception): ... class InvalidToken(Exception): def __str__(self) -> str: return 'Invalid token received, please try relog' + + +class APIException(Exception): + def __init__(self, detail: str = 'Unknown error', status: int = 400): + self.detail = _(detail) + self.status = status + + def __str__(self) -> str: + return f'{self.detail}' @@ -50,5 +50,20 @@ PATH_PREFETCH_MAP = { *_gen_only('payment_plan', 'uid', 'current_token_balance'), *_gen_only('business_account__parent_company__user__payment_plan', 'uid', 'current_token_balance') ) + }, + '/public/api-key': { + 'select': ( + 'host_account', + 'business_account', + ), + }, + '/public//': { + 'select': ('payment_plan',), + }, + '/public/openai/chat/completions': { + 'select': ('payment_plan',), + }, + '/public/openai/v1/chat/completions': { + 'select': ('payment_plan',), } } @@ -0,0 +1,22 @@ +import orjson +from ninja.parser import Parser + + +class ORJSONParser(Parser): + def parse_body(self, request): + return orjson.loads(request.body) + + +class MultiContentTypeParser(Parser): + def parse_body(self, request): + if request.content_type == 'application/json': + return orjson.loads(request.body) + elif request.content_type == 'multipart/form-data': + data = {} + for k,v in request.POST.dict().items(): + try: + data[k] =orjson.loads(v) + except Exception: + data[k] = v + data['file'] = request.FILES['file'] if request.FILES else None + return data @@ -1,5 +0,0 @@ -from django.urls import include, path - -urlpatterns = [ - path('public/', include('tools.public_api.urls')), -] @@ -333,6 +333,10 @@ FAL_API_KEY = env.str('FAL_API_KEY', 'defaultapikey') YANDEX_CLOUD_API_KEY = env.str('YANDEX_CLOUD_API_KEY', 'defaultapikey') YANDEX_CLOUD_ID = env.str('YANDEX_CLOUD_ID', 'defaultapikey') +OPENAI_PROXY_HOST = env.str('OPENAI_PROXY_HOST', 'neuron-proxy:8080') +UPSCALE_MULTIPLIER_HOST = env.str('UPSCALE_MULTIPLIER_HOST', 'packet:8080') + +DATA_UPLOAD_MAX_MEMORY_SIZE = 50 * 1024 * 1024 MAX_UPLOAD_SIZE_PER_MODEL = { 'raifgpt': 50, 'default': 8, @@ -13,16 +13,20 @@ from ninja import NinjaAPI from authentication.exceptions import ( InvalidPassword, InvalidToken, - InvalidUsername, + InvalidUsername, APIException, ) -from backend.public import urlpatterns as public_urlpatterns +from backend.parsers import ORJSONParser, MultiContentTypeParser -api = NinjaAPI(title='AIR', version='3.0.0') -api_debug = NinjaAPI(title='AIR API DEBUG', version='0.0.1', docs_url=None) +api = NinjaAPI(title='AIR', version='3.0.0', urls_namespace='air', parser=ORJSONParser()) +public_api = NinjaAPI(title='Public API', version='3.0.0', urls_namespace='public', parser=MultiContentTypeParser()) +api_debug = NinjaAPI(title='AIR API DEBUG', version='0.0.1', docs_url=None, urls_namespace='debug', parser=ORJSONParser()) api.add_router('ai/', 'ml_model.routes.v1.router') api.add_router('users/', 'users.routes.v1.router') api.add_router('chats/', 'tools.chats.routes.v3.router') api.add_router('media/', 'tools.media.routes.v3.router') + +public_api.add_router('', 'tools.public_api.routes.v3.router') + api_debug.add_router('auth/', 'authentication.routes.v1.router') api_debug.add_router('payments/', 'payments.routes.v1.router') @@ -49,6 +53,12 @@ def invalid_username_error_handler(request, exc: InvalidUsername): return api.create_response(request, {'message': _('Wrong username')}, status=401) +@api.exception_handler(APIException) +@public_api.exception_handler(APIException) +def api_exception(request, exc: APIException): + return api.create_response(request, {'message': str(exc)}, status=exc.status) + + def healthz_status(request): return JsonResponse({'status': 'ok'}, status=200) @@ -67,16 +77,15 @@ urlpatterns = [ SpectacularAPIView.as_view(urlconf=['backend.public']), name='schema-public', ), - path('api/v1/public/', include('tools.public_api.urls')), path( 'api/v1/public-view/', SpectacularSwaggerView.as_view(url_name='schema-public'), ), path('api/v1/api/', api.urls), path('api/v1/', api_debug.urls), + path('public/', public_api.urls) ] -urlpatterns += public_urlpatterns urlpatterns += static(settings.STATIC_URL, document_root=settings.STATIC_ROOT) if settings.DEBUG: @@ -1401,3 +1401,6 @@ msgstr "Плательщик не существует" msgid "The request must not be empty" msgstr "Запрос не должен быть пустым" + +msgid "Unknown error" +msgstr "Неизвестная ошибка" @@ -0,0 +1,97 @@ +from datetime import datetime, timedelta +from typing import Optional, Dict, Any, Union, List + +from django.conf import settings +from django.utils.translation import gettext_lazy as _ +from ninja import Schema, UploadedFile +from pydantic import UUID4, field_serializer, model_validator, Field + + +class BaseMessageSchema(Schema): + @model_validator(mode='after') + @classmethod + def check_file_size(cls, values): + file = getattr(values, 'file', None) + if isinstance(file, str): + return values + info = getattr(values, 'info', {}) or {} + version = info.get('inference', 'default') + + max_mb_size = settings.MAX_UPLOAD_SIZE_PER_MODEL.get( + version, + settings.MAX_UPLOAD_SIZE_PER_MODEL['default'] + ) + + if file and hasattr(file, 'size') and file.size > (max_mb_size << 10 << 10): + raise ValueError(_(f"The file size cannot exceed {max_mb_size} MB")) + + return values + + @field_serializer('content', check_fields=False) + def serialize_content(self, content, _info): + if content: + return content.replace('\\n', '\n') + return content + + +class InputMessageSchema(BaseMessageSchema): + content: str + file: Optional[Union[UploadedFile, str]] = None + info: Optional[Dict[str, Any]] = dict() + + +class OutputMessageSchema(InputMessageSchema): + uid: Optional[UUID4] = None + content: Optional[str] = None + file: Optional[Union[UploadedFile, str]] = None + from_model: Optional[bool] = None + model: Optional[str] = None + created_at: Optional[datetime] = None + elapsed_time: Optional[timedelta] = None + is_favourite: Optional[bool] = None + is_sent: Optional[bool] = None + info: Optional[Dict[str, Any]] = dict() + + @field_serializer('elapsed_time') + def serialize_elapsed_time(self, elapsed_time): + if elapsed_time is not None: + return elapsed_time.total_seconds() + return elapsed_time + + +class OpenaiInputMessageSchema(Schema): + model: str + messages: List[Dict[str, Any]] + stream: bool = False + + class Config: + extra = 'allow' + + +class OpenaiOutputMessageSchema(Schema): + id: UUID4 + object: str + created: Union[int, datetime] + model: str + choices: List[Dict[str, Any]] + service_tier: str + system_fingerprint: UUID4 + usage: Optional[Dict[str, Any]] = None + + + @field_serializer('created') + def serialize_datetime(self, created, _info): + if isinstance(created, str): + return int(datetime.fromisoformat(created.replace('Z', '+00:00')).timestamp()) + return int(created.timestamp()) + + +class OpenaiErrorDetailSchema(Schema): + message: str + type: str + param: Optional[str] = None + code: Optional[str] = None + + +class OpenaiErrorSchema(Schema): + error: OpenaiErrorDetailSchema @@ -2,19 +2,19 @@ from uuid import UUID from ninja import Router -from authentication.security import SyncAuthBearer +from authentication.security import AsyncAuthBearer from ml_model.schemas import InferenceSchema, NeuronModelSchema from ml_model.services.inference import InferenceService from ml_model.services.neuron_model import NeuronModelService -router = Router(auth=SyncAuthBearer(), tags=['ai']) +router = Router(auth=AsyncAuthBearer(), tags=['ai']) @router.get('model/{slug}/', response=NeuronModelSchema) -def get_model(request, slug: str): - return NeuronModelService.get_by_slug(slug=slug) +async def get_model(request, slug: str): + return await NeuronModelService.get_by_slug(slug=slug) @router.get('inferences/{id}/', response=InferenceSchema) -def get_inference(request, id: UUID): - return InferenceService.get_by_id(id=id) +async def get_inference(request, id: UUID): + return await InferenceService.get_by_id(id=id) @@ -50,8 +50,8 @@ class InferenceService: self.user = user @classmethod - def get_by_id(cls, id: UUID): - return Inference.objects.prefetch_related( + async def get_by_id(cls, id: UUID): + return await Inference.objects.prefetch_related( Prefetch( 'inference_parameters', queryset=OverridenParameter.objects.filter(parameter__hidden=False) ), @@ -62,11 +62,12 @@ class InferenceService: 'deployment__deployment_inputs', Prefetch('deployment__deployment_parameters', queryset=Parameter.objects.filter(hidden=False)), 'deployment__deployment_payment_rules', - ).get(id=id) + 'tags' + ).aget(id=id) @classmethod - def get_by_slug(cls, slug: str): - return Inference.objects.prefetch_related( + async def get_by_slug(cls, slug: str): + return await Inference.objects.prefetch_related( Prefetch( 'inference_parameters', queryset=OverridenParameter.objects.filter(parameter__hidden=False) ), @@ -77,7 +78,8 @@ class InferenceService: 'deployment__deployment_inputs', Prefetch('deployment__deployment_parameters', queryset=Parameter.objects.filter(hidden=False)), 'deployment__deployment_payment_rules', - ).get(slug=slug) + 'tags' + ).aget(slug=slug) def run( self, @@ -134,7 +136,7 @@ class InferenceService: break if not file_extension: raise UnknownFileException - if file_extension in ('png', 'jpg', 'jpeg'): + if file_extension in ('png', 'jpg', 'jpeg', 'webp'): file = BytesIO() normalized_image = ImageModule.open(file_buf) normalized_image.save(file, format='jpeg' if file_extension == 'jpg' else file_extension) @@ -20,8 +20,8 @@ class NeuronModelService: return NeuronModel.objects.get(uid=id) @classmethod - def get_by_slug(cls, slug: str) -> NeuronModel: - return NeuronModel.objects.prefetch_related( + async def get_by_slug(cls, slug: str) -> NeuronModel: + return await NeuronModel.objects.prefetch_related( Prefetch('inferences', queryset=Inference.objects.order_by('inferences_models__order')), Prefetch('inferences__tags', queryset=Tag.objects.order_by('tags_inferences__order')), - ).get(slug=slug) + ).aget(slug=slug) @@ -1,5 +1,7 @@ from django.utils.translation import gettext as _ +from authentication.exceptions import APIException + class GenerationException(Exception): def __str__(self): @@ -43,4 +45,10 @@ class FileExtensionNotSupported(Exception): class UnknownFileException(Exception): def __str__(self): - return _('Unknown file format') \ No newline at end of file + return _('Unknown file format') + + +class EnabledInferenceNotFound(APIException): + def __init__(self) -> None: + self.detail = 'Enabled inference not found in model' + self.status = 403 @@ -0,0 +1,367 @@ +import base64 +import orjson +import logging +import uuid +import time +from io import BytesIO + +import filetype + +from django.core.files.uploadedfile import InMemoryUploadedFile +from django.db.models import Q +from django.http import StreamingHttpResponse, JsonResponse +from django.utils.translation import gettext as _ +from django.core.cache import cache +from typing import List, Any, Dict + +from ninja import Router + +from authentication.exceptions import APIException +from authentication.security import AsyncAuthBearer +from messages.models import Message +from messages.schemas import OutputMessageSchema, InputMessageSchema, OpenaiInputMessageSchema, \ + OpenaiOutputMessageSchema, OpenaiErrorDetailSchema, OpenaiErrorSchema +from ml_model.exceptions import InferenceDisabled, EnabledInferenceNotFound +from ml_model.models import NeuronModel + +from ml_model.schemas import NeuronModelsSchema +from ml_model.services.neuron_model import NeuronModelService +from ml_model.tasks import run_inference +from payments.exceptions.insufficient_balance import InsufficientBalance +from tools.public_api.exceptions import APIKeyLimitExceeded, ModelBlockException + +from tools.public_api.models import APIStore, APIKey +from tools.public_api.schemas import APIKeySchema, APIKeyCreateSchema, APIKeyUpdateSchema, APIKeyDeleteSchema +from tools.public_api.security import APIKeyAuthentication +from tools.public_api.services import APIKeyService + +logger = logging.getLogger(__name__) + +router = Router(auth=AsyncAuthBearer(), tags=['public']) + + +@router.post('api-key', tags=['public/api-key'], response=APIKeySchema, summary='Create the API key') +async def create_api_key(request, payload: APIKeyCreateSchema): + try: + api_key = await APIKeyService(request.auth).create(payload=dict(payload)) + return api_key + except Exception as exc: + logger.exception(exc) + raise APIException from exc + + +@router.get('api-key', tags=['public/api-key'], response=List[APIKeySchema], summary='Get all API user keys') +async def list_api_keys(request): + try: + api_keys = [ + api_key + async for api_key in APIKey.objects.filter(user=request.auth, is_deleted=False) + ] + return api_keys + except Exception as exc: + logger.exception(exc) + raise APIException from exc + + +@router.patch('api-key', tags=['public/api-key'], response=APIKeySchema, summary='Update the API key') +async def update_api_key_by_name(request, payload: APIKeyUpdateSchema): + try: + api_key = await APIKeyService(request.auth).update(**dict(payload)) + return api_key + except Exception as exc: + logger.exception(exc) + raise APIException from exc + + +@router.delete('api-key', tags=['public/api-key'], response=Dict[str, bool], summary='Delete the API key') +async def delete_api_key_by_name(request, payload: APIKeyDeleteSchema): + try: + result = await APIKeyService(request.auth).delete(**dict(payload)) + return {'result': result} + except Exception as exc: + logger.exception(exc) + raise APIException from exc + + +@router.get( + path='/{content_type}', + tags=['public/{content_type}'], + response=List[NeuronModelsSchema], + summary='List available generative Models for output content type: text, image, audio, video, or code.' +) +async def get_models_by_ct(request, content_type: str): + try: + categories = { + 'text': 'chat-bots', + 'image': 'images', + 'audio': 'audio', + 'video': 'videos', + 'code': 'code' + } + return [ + model + async for model in NeuronModelService.list_all(types=[categories[content_type]]) + ] + except Exception as exc: + logger.exception(exc) + raise APIException from exc + + +@router.post( + auth=APIKeyAuthentication(), + path='/{content_type}/{model}', + tags=['public/{content_type}/{model}'], + response=OutputMessageSchema | Any, + summary='Create new content. Type of content depends on model output content type.' +) +async def base_generate(request, payload: InputMessageSchema, content_type: str, model: str, stream: bool=False): + user, key = request.auth + balance = user.balance + if key.token_limit is not None and key.token_limit < 1: + return APIKeyLimitExceeded + model = await NeuronModelService.get_by_slug(slug=model) + store, created = await APIStore.objects.aget_or_create(user=user, model=model) + if not model.enabled: + raise ModelBlockException + message_data = {k: v for k, v in dict(payload).items() if v is not None} + info = message_data.pop('info') + inference_slug = info.pop('inference') + if not any( + [ + inference.slug == inference_slug and inference.enabled + async for inference in model.inferences.all() + ] + ): + raise EnabledInferenceNotFound + input_message = await Message.objects.acreate( + **message_data, content_object=store, info=info, from_model=False + ) + + task_result = run_inference.delay( + user_id=user.uid, inference_slug=inference_slug, input_message_id=input_message.uid + ) + output_slot_id = task_result.get() + output_message = await Message.objects.aget(uid=output_slot_id) + + async def message_stream(): + yield f'id: {output_slot_id}\nevent: start\ndata: [START]\n\n' + + try: + cache_key = f'apistores:{store.uid}' + content = await cache.aget(cache_key, default=[]) + while cache.has_key(cache_key): + chunk = (await cache.aget(cache_key, default=[]))[len(content):] + if chunk: + content += chunk + yield (f'id: {output_slot_id}\nevent: output\ndata: ' + f'{"".join(map(lambda x: x['content'], chunk)).replace("\n", "\\n")}\n\n') + yield f'id: {output_slot_id}\nevent: done\ndata: [DONE]\n\n' + except Exception as exc: + input_message.is_sent = False + await input_message.asave() + if not isinstance(exc, (InsufficientBalance, InferenceDisabled)): + logger.exception(exc) + yield f'id: {output_slot_id}\nevent: error\ndata: {exc}\n\n' + + if key.token_limit is not None: + user_after = user + key.token_limit -= balance - user_after.balance + await key.asave() + + await output_message.arefresh_from_db() + output_message.from_public_api = True + await output_message.asave() + + if stream: + return StreamingHttpResponse(message_stream(), content_type='text/event-stream') + + async for chunk in message_stream(): + pass + + return output_message + + +@router.post( + auth=APIKeyAuthentication(), + path='/openai/chat/completions', + tags=['public/openai/chat/completions'], + response=OpenaiOutputMessageSchema | Any, + summary='Create new content using OpenAI SDK.', +) +@router.post( + auth=APIKeyAuthentication(), + path='openai/v1/chat/completions', + tags=['public/openai/v1/chat/completions'], + response=OpenaiOutputMessageSchema | Any, + summary='Create new content using OpenAI SDK.' +) +async def openai_compatible_generate(request, payload: OpenaiInputMessageSchema): + data = payload.dict() + stream = data.pop('stream') + try: + version = data.pop('model') + except KeyError: + return JsonResponse(OpenaiErrorSchema( + error=OpenaiErrorDetailSchema( + message=_('You must provide a model parameter'), + type='invalid_request_error', + ) + ).dict(), status=400) + content_lines = [] + file = None + messages = data.pop('messages', []) + if not messages: + return JsonResponse(OpenaiErrorSchema( + error=OpenaiErrorDetailSchema( + message=_("Missing required parameter: 'messages'"), + type='invalid_request_error', + param='messages', + code='missing_required_parameter' + ) + ).dict(), status=400) + for m in messages: + role = m['role'].capitalize() + msg_content = m.get('content', []) + if isinstance(msg_content, str): + content_lines.append(f'[{role}] {msg_content}') + else: + for c in msg_content: + if c['type'] == 'text': + content_lines.append(f'[{role}] {c["text"]}') + elif c['type'] == 'image_url': + data_url = c['image_url']['url'] + encoded = data_url.split('base64')[-1] + buf = BytesIO(base64.b64decode(encoded)) + kind = filetype.guess(buf.read(20)) + buf.seek(0) + mime = kind.mime if kind else 'application/octet-stream' + file = InMemoryUploadedFile( + buf, + field_name='file', + name=f'api-file.png', + content_type=mime, + size=buf.getbuffer().nbytes, + charset=None, + ) + content = '\n'.join(content_lines) + data['info'] = {key: data.pop(key) for key in data.copy().keys()} + data['info']['inference'] = version + + try: + model = await NeuronModel.objects.filter( + Q(inferences__slug=data['info']['inference']) | Q(slug=data['info']['inference']) + ).afirst() + if 'images' in model.types: + return JsonResponse(OpenaiErrorSchema( + error=OpenaiErrorDetailSchema( + message=_('You cannot use a non-text model'), + type='invalid_model_type', + ) + ).dict(), status=400) + except NeuronModel.DoesNotExist: + return JsonResponse(OpenaiErrorSchema( + error=OpenaiErrorDetailSchema( + message=_('Model not found'), + type='invalid_request_error', + ) + ).dict(), status=400) + data['content'] = content + data['file'] = file + + result = await base_generate(request, data, model.types, model.slug, stream) + + if stream: + async def openai_stream(): + yield f"data: {orjson.dumps( + { + 'id': str(uuid.uuid4()), + 'object': 'chat.completion.chunk', + 'created': int(time.time()), + 'model': version, + 'service_tier': 'auto', + 'system_fingerprint': str(uuid.uuid4()), + 'choices': [{'index': 0, 'delta': {'role': 'assistant'}, 'finish_reason': None}]} + ).decode('utf-8')}\n\n" + + async for chunk_bytes in result.streaming_content: + chunk = chunk_bytes.decode('utf-8') if isinstance(chunk_bytes, bytes) else chunk_bytes + + if chunk.startswith('id: ') and 'event: output' in chunk: + lines = chunk.strip().split('\n') + data_line = None + for line in lines: + if line.startswith('data: '): + data_line = line[6:] + break + + if data_line and data_line != '[START]': + chunk_data = { + 'id': str(uuid.uuid4()), + 'object': 'chat.completion.chunk', + 'created': int(time.time()), + 'model': version, + 'service_tier': 'auto', + 'system_fingerprint': str(uuid.uuid4()), + 'choices': [{ + 'index': 0, + 'delta': {'content': data_line}, + 'finish_reason': None + }] + } + yield f"data: {orjson.dumps(chunk_data).decode('utf-8')}\n\n" + elif 'event: done' in chunk: + final_data = { + 'id': str(uuid.uuid4()), + 'object': 'chat.completion.chunk', + 'created': int(time.time()), + 'model': version, + 'service_tier': 'auto', + 'system_fingerprint': str(uuid.uuid4()), + 'choices': [{ + 'index': 0, + 'delta': {}, + 'finish_reason': 'stop' + }] + } + yield f"data: {orjson.dumps(final_data).decode('utf-8')}\n\n" + yield "data: [DONE]\n\n" + break + elif 'event: error' in chunk: + error_data = { + 'id': str(uuid.uuid4()), + 'object': 'chat.completion.chunk', + 'created': int(time.time()), + 'model': version, + 'service_tier': 'auto', + 'system_fingerprint': str(uuid.uuid4()), + 'choices': [{ + 'index': 0, + 'delta': {}, + 'finish_reason': 'stop' + }], + 'error': { + 'message': 'Generation failed', + 'type': 'generation_error' + } + } + yield f"data: {orjson.dumps(error_data).decode('utf-8')}\n\n" + yield "data: [DONE]\n\n" + break + + return StreamingHttpResponse(openai_stream(), content_type="text/event-stream") + + return OpenaiOutputMessageSchema( + id=uuid.uuid4(), + object='chat.completion', + created=result.created_at, + model=version, + choices=[ + { + 'index': 0, + 'message': {'role': 'assistant', 'content': result.content}, + 'finish_reason': 'stop', + } + ], + service_tier='auto', + system_fingerprint=uuid.uuid4(), + ) @@ -1 +0,0 @@ -from .api_key import APIKeySelector @@ -1,26 +0,0 @@ -from core.selector import BaseSelector -from tools.public_api.exceptions import APIKeyNotFound -from tools.public_api.models import APIKey -from tools.public_api.serializers import APIKeyResultSerializer - - -class APIKeySelector(BaseSelector): - def list(self, serialize: bool = False): - api_keys = APIKey.objects.filter(user=self.user, is_deleted=False) - if serialize: - return APIKeyResultSerializer(api_keys, many=True) - return api_keys - - def get_by_name(self, name: str, serialize: bool = False): - api_key = APIKey.objects.get(user=self.user, name=name, is_deleted=False) - if serialize: - return APIKeyResultSerializer(api_key) - return api_key - - @classmethod - def get_user_by_key(cls, key_value: str): - try: - api_key = APIKey.objects.get(key=key_value) - return api_key.user - except APIKey.DoesNotExist: - raise APIKeyNotFound @@ -1 +1 @@ -from .api_key import APIKeyService +from .api_key import APIKeyService \ No newline at end of file @@ -1,29 +1,23 @@ from datetime import date from decimal import Decimal +from asgiref.sync import sync_to_async from django.contrib.admin.models import ADDITION, DELETION, LogEntry from django.contrib.contenttypes.models import ContentType -from authentication.selectors.account_status_selector import ( - AccountStatusSelector, -) from core.service import BaseService from tools.public_api.models import APIKey -from tools.public_api.selectors.api_key import APIKeySelector -from tools.public_api.serializers import APIKeyResultSerializer class APIKeyService(BaseService): - def create(self, payload: dict, serialize: bool = False) -> APIKey | APIKeyResultSerializer: - if not ( - AccountStatusSelector(self.user).is_business_host() - or AccountStatusSelector(self.user).is_admin() - ): + async def create(self, payload: dict) -> APIKey: + if not self.user.account_type in ('business_host', 'regular', 'business_admin'): raise Exception('Can not create API key from business sub-account.') - api_key = APIKey.objects.create(user=self.user, **payload) - LogEntry.objects.log_action( + api_key = await APIKey.objects.acreate(user=self.user, **payload) + content_type = await sync_to_async(ContentType.objects.get_for_model)(api_key) + await sync_to_async(LogEntry.objects.log_action)( self.user.pk, - ContentType.objects.get_for_model(api_key).pk, + content_type.pk, api_key.pk, str(api_key), ADDITION, @@ -36,36 +30,32 @@ class APIKeyService(BaseService): } ], ) - if serialize: - return APIKeyResultSerializer(api_key) return api_key - def update( + async def update( self, - key_name: str, + name: str, new_name: str | None = None, expires_at: date | None = None, token_limit: Decimal | None = None, - serialize: bool = False, - ) -> APIKey | APIKeyResultSerializer: - api_key = APIKeySelector(self.user).get_by_name(name=key_name) + ) -> APIKey : + api_key = await APIKey.objects.aget(name=name) if new_name is not None: api_key.name = new_name if expires_at is not None: api_key.expires_at = expires_at api_key.token_limit = token_limit - api_key.save() + await api_key.asave() - if serialize: - return APIKeyResultSerializer(api_key) return api_key - def delete(self, key_name): - api_key = APIKeySelector(self.user).get_by_name(name=key_name) - LogEntry.objects.log_action( + async def delete(self, name): + api_key = await APIKey.objects.aget(name=name) + content_type = await sync_to_async(ContentType.objects.get_for_model)(api_key) + sync_to_async(LogEntry.objects.log_action)( self.user.pk, - ContentType.objects.get_for_model(api_key).pk, + content_type.pk, api_key.pk, str(api_key), DELETION, @@ -79,5 +69,5 @@ class APIKeyService(BaseService): ], ) api_key.is_deleted = True - api_key.save() + await api_key.asave() return True @@ -1,2 +0,0 @@ -from .api_key import APIKeyView -from .user import UserInfoAPIView @@ -1,75 +0,0 @@ -import logging - -from drf_spectacular.utils import extend_schema -from rest_framework import status -from rest_framework.permissions import IsAuthenticated -from rest_framework.response import Response -from rest_framework.views import APIView - -from tools.public_api.selectors.api_key import APIKeySelector -from tools.public_api.serializers import ( - APIKeyCreateSerializer, - APIKeyDeleteSerializer, - APIKeyResultSerializer, - APIKeyUpdateSerializer, -) -from tools.public_api.services.api_key import APIKeyService - -logger = logging.getLogger(__name__) - - -class APIKeyView(APIView): - permission_classes = (IsAuthenticated,) - - @extend_schema(request=APIKeyCreateSerializer, responses={201: APIKeyResultSerializer}) - def post(self, request, *args, **kwargs): - """Create new API Key for user.""" - try: - serializer = APIKeyCreateSerializer(data=request.data) - serializer.is_valid(raise_exception=True) - api_key = APIKeyService(self.request.user).create(payload=serializer.validated_data) - result = APIKeyResultSerializer(api_key) - - return Response(result.data, status=status.HTTP_201_CREATED) - - except Exception as err: - logger.exception(err) - return Response({'detail': str(err)}, status=status.HTTP_400_BAD_REQUEST) - - def get(self, request, *args, **kwargs): - """List user API Keys.""" - try: - keys_selector = APIKeySelector(self.request.user) - api_keys = keys_selector.list(serialize=True) - return Response(api_keys.data) - except Exception as err: - logger.exception(err) - return Response({'detail': str(err)}, status=status.HTTP_400_BAD_REQUEST) - - @extend_schema(request=APIKeyUpdateSerializer, responses={201: APIKeyResultSerializer}) - def patch(self, request, *args, **kwargs): - """Edit user's API Key.""" - try: - serializer = APIKeyUpdateSerializer(data=request.data) - serializer.is_valid(raise_exception=True) - api_key = APIKeyService(self.request.user).update( - key_name=serializer.validated_data.get('name'), - new_name=serializer.validated_data.get('new_name'), - expires_at=serializer.validated_data.get('expires_at'), - token_limit=serializer.validated_data.get('token_limit'), - serialize=True, - ) - return Response(api_key.data, status=status.HTTP_200_OK) - except Exception as err: - return Response({'detail': str(err)}, status=status.HTTP_400_BAD_REQUEST) - - @extend_schema(request=APIKeyDeleteSerializer, responses={204: bool}) - def delete(self, request, *args, **kwargs): - """Delete user's API Key.""" - try: - serializer = APIKeyDeleteSerializer(data=request.data) - serializer.is_valid(raise_exception=True) - result = APIKeyService(self.request.user).delete(key_name=serializer.validated_data['name']) - return Response(data={'result': result}, status=status.HTTP_200_OK) - except Exception as err: - return Response({'detail': str(err)}, status=status.HTTP_400_BAD_REQUEST) @@ -1,119 +0,0 @@ -import logging - -from django.core.cache import cache -from django.http import StreamingHttpResponse -from django.utils.translation import gettext_lazy as _ -from rest_framework.response import Response -from rest_framework.status import ( - HTTP_403_FORBIDDEN, -) -from rest_framework.views import APIView - -from messages.models import Message -from messages.serializers import MessageSerializer -from ml_model.exceptions import InferenceDisabled -from ml_model.models import NeuronModel -from ml_model.schemas import NeuronModelsSchema -from ml_model.services.neuron_model import NeuronModelService -from ml_model.tasks import run_inference -from payments.exceptions.insufficient_balance import InsufficientBalance -from tools.public_api.models import APIKey, APIStore -from tools.public_api.permissions import HasAPIKey -from tools.public_api.selectors.api_key import APIKeySelector - -logger = logging.getLogger(__name__) - - -class ListModelsByContentTypeAPIView(APIView): - permission_classes = (HasAPIKey,) - - def get(self, request, content_type: str, *args, **kwargs): - """List available generative Models for output content type: - text, image, audio, video, or code.""" - return Response( - NeuronModelsSchema.from_orm(model).model_dump_json() - for model in NeuronModelService.list_all(types=[content_type]) - ) - - -class BaseGenerationView(APIView): - permission_classes = (HasAPIKey,) - - def post(self, request, model, *args, **kwargs): - """Create new content. Type of content depends on model output content type: - text, image, audio, video, or code.""" - user = APIKeySelector.get_user_by_key(key_value=request.headers.get('Authorization', '').split()[-1]) - balance = user.balance - key = APIKey.objects.get(key=request.headers.get('Authorization', '')) - if key.token_limit is not None and key.token_limit < 1: - return Response({'detail': _('Key limit exceeded')}, HTTP_403_FORBIDDEN) - model: NeuronModel = NeuronModelService.get_by_slug(slug=model) - store, created = APIStore.objects.get_or_create(user=user, model=model) - - if not model.enabled: - return Response( - {'detail': _('Model is blocked by outdating or temporary block, please retry later')}, - status=HTTP_403_FORBIDDEN, - ) - serializer = MessageSerializer(data=request.data) - try: - serializer.is_valid(raise_exception=True) - info = serializer.validated_data.pop('info') - inference_slug = info.pop('inference') - if not any( - [ - inference.slug == inference_slug and inference.enabled - for inference in model.inferences.all() - ] - ): - raise Exception(_('Enabled inference not found in model')) - input_message = Message.objects.create( - **serializer.validated_data, content_object=store, info=info, from_model=False - ) - except Exception as exc: - return Response({'detail': str(exc)}, status=400) - - task_result = run_inference.delay( - user_id=key.user.uid, inference_slug=inference_slug, input_message_id=input_message.uid - ) - output_slot_id = task_result.get() - output_message = Message.objects.get(uid=output_slot_id) - - async def message_stream(): - yield f'id: {output_slot_id}\nevent: start\ndata: [START]\n\n' - - try: - cache_key = f'messages:{output_slot_id}' - content = await cache.aget(cache_key, default=[]) - while cache.has_key(cache_key): - chunk = (await cache.aget(cache_key, default=[]))[len(content) :] - if chunk: - content += chunk - yield f'id: {output_slot_id}\nevent: output\ndata: {"".join(chunk).replace("\n", "\\n")}\n\n' - - yield f'id: {output_slot_id}\nevent: done\ndata: [DONE]\n\n' - except Exception as exc: - input_message.is_sent = False - await input_message.asave() - if not isinstance(exc, (InsufficientBalance, InferenceDisabled)): - logger.exception(exc) - yield f'id: {output_slot_id}\nevent: error\ndata: {exc}\n\n' - - if key.token_limit is not None: - user_after = APIKeySelector.get_user_by_key( - key_value=request.headers.get('Authorization', '') - ) - key.token_limit -= balance - user_after.balance - await key.asave() - - await output_message.arefresh_from_db() - output_message.from_public_api = True - await output_message.asave() - - if request.query_params.get('stream', 'false') == 'true': - return StreamingHttpResponse(message_stream(), content_type='text/event-stream') - - for chunk in message_stream(): - pass - - return Response(MessageSerializer(output_message).data) @@ -1,23 +0,0 @@ -import logging - -from rest_framework.response import Response -from rest_framework.views import APIView -from tools.public_api.permissions import HasAPIKey -from tools.public_api.selectors.api_key import APIKeySelector - -from authentication.serializers import UserDetailSerializer - -logger = logging.getLogger(__name__) - - -class UserInfoAPIView(APIView): - """Testing class for availability and correct functioning of API Token.""" - - permission_classes = (HasAPIKey,) - - def get(self, request, *args, **kwargs): - """Fetch user info by API Token.""" - user = APIKeySelector.get_user_by_key(key_value=request.headers.get('Authorization', '')) - return Response( - {k: v for k, v in UserDetailSerializer(user).data.items() if k not in ('payment_plan', 'token')} - ) @@ -1,14 +1,20 @@ -from django.utils.translation import gettext_lazy as _ -from rest_framework.exceptions import APIException +from authentication.exceptions import APIException -class LowTokenLimit(APIException): - status_code = 403 - default_detail = _('Upgrade token limit on your api-key') - default_code = 'low_limit' +class APIKeyNotFound(APIException): + def __init__(self) -> None: + self.detail = 'API Key not found' + self.status = 400 + super().__init__(self.detail, self.status) -class APIKeyNotFound(APIException): - status_code = 400 - default_detail = _('API Key not found') - default_code = 'api_key_not_found' +class APIKeyLimitExceeded(APIException): + def __init__(self) -> None: + self.detail = 'Key limit exceeded' + self.status = 403 + + +class ModelBlockException(APIException): + def __init__(self) -> None: + self.detail = 'Model is blocked by outdating or temporary block, please retry later' + self.status = 403 \ No newline at end of file @@ -1,19 +0,0 @@ -from datetime import date - -from rest_framework import permissions -from rest_framework.exceptions import PermissionDenied - -from tools.public_api.models import APIKey - - -class HasAPIKey(permissions.BasePermission): - def has_permission(self, request, view): - api_key_value = request.headers.get('Authorization', '') - if api_key_value.startswith('Bearer'): - api_key_value = api_key_value[len('Bearer') + 1 :] - api_key: APIKey = APIKey.objects.get_or_none(key=api_key_value) - if not api_key: - raise PermissionDenied('No API Key in Authorization header') - if api_key.expires_at and api_key.expires_at < date.today(): - raise PermissionDenied('API Key is expired') - return api_key @@ -0,0 +1,35 @@ +from datetime import datetime +from typing import Optional + +from ninja import ModelSchema, Schema +from pydantic import constr, condecimal + +from tools.public_api.models import APIKey + + +class APIKeySchema(ModelSchema): + class Meta: + model = APIKey + fields = ( + 'created_at', + 'name', + 'key', + 'expires_at', + 'token_limit', + ) + + +class APIKeyCreateSchema(Schema): + name: constr(max_length=255) + expires_at: datetime + + +class APIKeyUpdateSchema(Schema): + name: constr(max_length=255) + new_name: Optional[constr(max_length=255)] = None + expires_at: Optional[datetime] = None + token_limit: Optional[condecimal(decimal_places=2, max_digits=50)] = None + + +class APIKeyDeleteSchema(Schema): + name: constr(max_length=255) @@ -0,0 +1,37 @@ +from datetime import date + +from django.urls import resolve +from django.utils.translation import gettext as _ + +from ninja.errors import HttpError +from ninja.security import APIKeyHeader + +from authentication.mapper import PATH_PREFETCH_MAP +from authentication.models import CustomUserModel +from tools.public_api.models import APIKey + + +class APIKeyAuthentication(APIKeyHeader): + param_name = 'Authorization' + + async def authenticate(self, request, key): + try: + mapper = PATH_PREFETCH_MAP[request.path] + except KeyError: + mapper = PATH_PREFETCH_MAP.get(f'/{resolve(request.path).route}', {}) + if (split_api_key := key.split())[0] == 'Bearer': + key = split_api_key[-1] + try: + api_key = await APIKey.objects.aget(key=key) + except APIKey.DoesNotExist: + raise HttpError(401, _('No API Key in Authorization header')) + if api_key.expires_at and api_key.expires_at < date.today(): + raise HttpError(401, _('API Key is expired')) + api_key_user = await CustomUserModel.objects.select_related( + *mapper.get('select', []) + ).prefetch_related( + *mapper.get('prefetch', []) + ).only( + *mapper.get('only', []) + ).filter(apikey__key=key).afirst() + return api_key_user, api_key @@ -1,46 +0,0 @@ -from rest_framework import serializers - -from authentication.serializers import UserDetailSerializer -from tools.public_api.models import APIKey - -__all__ = [ - 'APIKeyResultSerializer', - 'APIKeyCreateSerializer', - 'APIKeyUpdateSerializer', - 'APIKeyDeleteSerializer', -] - - -class APIKeyResultSerializer(serializers.ModelSerializer): - user = UserDetailSerializer() - - class Meta: - model = APIKey - fields = ( - 'created_at', - 'name', - 'key', - 'expires_at', - 'user', - 'token_limit', - ) - - -class APIKeyCreateSerializer(serializers.ModelSerializer): - class Meta: - model = APIKey - fields = ('name', 'expires_at') - - -class APIKeyUpdateSerializer(serializers.ModelSerializer): - new_name = serializers.CharField(required=False) - - class Meta: - model = APIKey - fields = ('name', 'new_name', 'expires_at', 'token_limit') - - -class APIKeyDeleteSerializer(serializers.ModelSerializer): - class Meta: - model = APIKey - fields = ('name',) @@ -1,3 +0,0 @@ -# from django.test import TestCase # Flake8 angery >:( - -# Create your tests here. @@ -1,11 +0,0 @@ -from django.urls import path - -from tools.public_api import views -from tools.public_api.views.base import BaseGenerationView, ListModelsByContentTypeAPIView - -urlpatterns = [ - path('api-key', views.APIKeyView.as_view()), - path('me', views.UserInfoAPIView.as_view()), - path('', ListModelsByContentTypeAPIView.as_view()), - path('/', BaseGenerationView.as_view()), -]