@@ -74,6 +74,7 @@ from ml_model.services.sora import Sora from ml_model.services.stablediffusion import Stablediffusion from ml_model.services.stablemusic import Stablemusic from ml_model.services.suno import Suno +from ml_model.services.text_test_model import Text_Test_Model from ml_model.services.upscaleai import Upscaleai from ml_model.services.veo import Veo from ml_model.services.vicuna import Vicuna @@ -0,0 +1,81 @@ +import hashlib +import time +import tiktoken + +from datetime import timedelta +from decimal import Decimal + +from django.core.cache import cache + +from messages.models import Message +from ml_model.services.base import SimpleService + + +class Text_Test_Model(SimpleService): + TOKENS_COST = { + 'input': Decimal('1000'), # per 1 million input-tokens + 'output': Decimal('2500'), # per 1 million output-tokens + } + + ENCODING = 'o200k_base' + + BASE_OUTPUT_MESSAGE = """ + Lorem ipsum dolor sit amet, consectetur adipiscing elit. Pellentesque ac metus ac dolor + mattis euismod in id eros. Phasellus sed ornare ligula, sit amet ullamcorper ante. Aenean + fermentum lorem sit amet tortor ultricies, id pulvinar nibh pulvinar. + """ + + def calculate_price(self, input_tokens: int, output_tokens: int) -> Decimal: + price = ( + input_tokens * self.TOKENS_COST['input'] / 1_000_000 + + output_tokens * self.TOKENS_COST['output'] / 1_000_000 + ) + return price.quantize(Decimal('0.01'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + ) + if save: + msg.save() + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + info = input_message.info.copy() + start_time = time.time() + input_tokens = self._get_cached_tokens(input_message.content) + output_tokens = self._get_cached_tokens(info.get('cm') or self.BASE_OUTPUT_MESSAGE) + ttft = info.get('ttft', 0.5) + tbt = info.get('tbt', 0.35) + time.sleep(ttft) + result = '' + for i, token in enumerate(output_tokens, start=1): + result += token + if i < len(output_tokens): + time.sleep(tbt) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, len(input_tokens), len(output_tokens)) + return self.save_results(result, process_time, save) + + @classmethod + def _tokenize(cls, text: str) -> list[str]: + encoding = tiktoken.get_encoding(cls.ENCODING) + return [encoding.decode([token]) for token in encoding.encode(text)] + + @classmethod + def _get_token_cache_key(cls, text: str) -> str: + return hashlib.sha256(f'text_test_model:tokens:{cls.ENCODING}:{text}'.encode('utf-8')).hexdigest() + + @classmethod + def _get_cached_tokens(cls, text: str) -> list[str]: + key = cls._get_token_cache_key(text) + cached = cache.get(key) + if cached is not None: + return cached + tokens = cls._tokenize(text) + cache.set(key, tokens) + return tokens + + @@ -0,0 +1,89 @@ +import json +from decimal import Decimal +from datetime import timedelta +from unittest.mock import patch + +from core.tests import BaseAuthorizedAPITest +from ml_model.models import ModelCategory, NeuronModel +from tools.chats.models import Chat + + +class TextTestModelAPITest(BaseAuthorizedAPITest): + INPUT_TEXT = ( + 'Lorem ipsum dolor sit amet, consectetur adipiscing elit. Pellentesque ac metus ac dolor mattis euismod in ' + 'id eros. Phasellus sed ornare ligula, sit amet ullamcorper ante. Aenean fermentum lorem sit amet tortor ' + 'ultricies, id pulvinar nibh pulvinar.' + ) + OUTPUT_TEXT = ( + 'Aliquam molestie orci nisl, eget rhoncus nisi varius non. Integer eleifend neque nisi, quis feugiat augue ' + 'malesuada eu. Mauris tincidunt augue id justo ultrices convallis. Nullam a lorem mauris. Duis faucibus est ' + 'mauris, id vestibulum tellus tempor rhoncus.' + ) + TTFT = 0.5 + TBT = 0.35 + + @classmethod + def setup_test_data(cls) -> None: + category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots') + cls.model = NeuronModel.objects.create( + title='Text Test Model', + slug='text_test_model', + category=category, + ) + cls.chat = Chat.objects.create(title='Text test chat', user=cls.user, model=cls.model) + + @property + def ENDPOINT(self) -> str: + return f'/api/v1/chats/{self.chat.uid}/messages/' + + def _request_data(self) -> dict: + return { + 'content': self.INPUT_TEXT, + 'info': json.dumps( + { + 'ttft': self.TTFT, + 'tbt': self.TBT, + 'cm': self.OUTPUT_TEXT, + } + ), + } + + @staticmethod + def _duration_from_api(value: str) -> Decimal: + hours, minutes, seconds = value.split(':') + total = timedelta( + hours=int(hours), + minutes=int(minutes), + seconds=float(seconds), + ).total_seconds() + return Decimal(str(total)) + + def test_unauthorized_status_code(self) -> None: + response = self.client.post(self.ENDPOINT, data=self._request_data()) + self.assertEqual(response.status_code, 401) + self.assertIn('detail', response.json()) + + @patch('ml_model.services.text_test_model.time.sleep') + def test_generation_time_is_not_more_than_21_seconds(self, _sleep_mock) -> None: + response = self.post(data=self._request_data()) + self.assertEqual(response.status_code, 201, response.json()) + + payload = response.json() + self.assertEqual(len(payload), 2) + self.assertEqual(payload[1]['content'], self.OUTPUT_TEXT) + + elapsed = self._duration_from_api(payload[1]['elapsed_time']) + self.assertLessEqual(elapsed, Decimal('21')) + + @patch('ml_model.services.text_test_model.time.sleep') + def test_billing_uses_rounded_price(self, _sleep_mock) -> None: + expected_charge = Decimal('0.20') + + balance_before = self.user.payment_plan.current_token_balance + response = self.post(data=self._request_data()) + self.assertEqual(response.status_code, 201, response.json()) + + self.user.payment_plan.refresh_from_db() + charged_amount = balance_before - self.user.payment_plan.current_token_balance + + self.assertEqual(charged_amount, expected_charge)