From 64901ecb0b23f355884d4be27f380927605f79f2 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 12:16:50 +0300 Subject: [PATCH 01/46] =?UTF-8?q?docs:=20=D0=B4=D0=BE=D0=B1=D0=B0=D0=B2?= =?UTF-8?q?=D0=B8=D1=82=D1=8C=20=D0=BF=D0=BB=D0=B0=D0=BD=20=D0=BF=D0=BE?= =?UTF-8?q?=D1=88=D0=B0=D0=B3=D0=BE=D0=B2=D0=BE=D0=B3=D0=BE=20=D1=80=D0=B5?= =?UTF-8?q?=D1=84=D0=B0=D0=BA=D1=82=D0=BE=D1=80=D0=B8=D0=BD=D0=B3=D0=B0=20?= =?UTF-8?q?=D0=BF=D1=80=D0=BE=D0=B5=D0=BA=D1=82=D0=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: aider (ollama/kimi-k2.6:latest-cloud) --- global_to_do.txt | 151 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 151 insertions(+) create mode 100644 global_to_do.txt diff --git a/global_to_do.txt b/global_to_do.txt new file mode 100644 index 0000000..6cd4622 --- /dev/null +++ b/global_to_do.txt @@ -0,0 +1,151 @@ +================================================================================ + ПЛАН РЕФАКТОРИНГА ASR FastAPI + "Одна задача за раз" +================================================================================ + +ПРИНЦИПЫ РАБОТЫ: + - Одна задача за раз. + - Жди подтверждения перед правками. + - Сначала описывается план, пользователь соглашается — потом выполняется. + - Не делать несколько изменений сразу. + - После каждого этапа требуется тестирование перед переходом к следующему. + +================================================================================ +ЭТАП 1. ФУНДАМЕНТ: КОНФИГУРАЦИЯ И ТОЧКА ВХОДА +================================================================================ + +Цель: устранить side-effects при импорте, валидировать окружение и централизовать +создание приложения. + +ЗАДАЧА 1.1 — config.py + - Миграция с голого os.getenv() на pydantic-settings (BaseSettings). + - Все переменные получают типизацию, значения по умолчанию и валидацию. + - Старые имена экспортируются для обратной совместимости с неизменяемыми + модулями ASR (Recognizer, Diarisation и т.д.). + - Добавить SettingsConfigDict с поддержкой .env файла. + +ЗАДАЧА 1.2 — utils/pre_start_init.py + - Удалить глобальный объект app = FastAPI(...). + - Оставить только инициализацию paths, глобальных defaultdict + (audio_overlap, audio_buffer, audio_to_asr, audio_duration, + ws_collected_asr_res, posted_and_downloaded_audio) и lifespan. + - Убрать импорт FastAPI из этого файла. + - Убрать любие импорты роутов или HTTP-объектов. + - Файл должен быть безопасен для импорта без побочных эффектов. + +ЗАДАЧА 1.3 — main.py + - Создание app = FastAPI(...) переносится сюда как явная фабрика. + - Подключение lifespan из utils.pre_start_init. + - Добавление CORS-middleware (CORSMiddleware). + - Перенос монтирования статики (/static) из routes/demo_page.py сюда. + - Точка входа if __name__ == '__main__': остается здесь. + +================================================================================ +ЭТАП 2. МАРШРУТИЗАЦИЯ: ОТКАЗ ОТ ГЛОБАЛЬНОГО app +================================================================================ + +Цель: разорвать жесткую связность роутов, убрать циклические импорты +и подготовить версионирование API. + +ЗАДАЧА 2.1 — routes/root.py, routes/is_alive.py, routes/post_ws.py + - Заменить импорт глобального app на создание APIRouter() в каждом файле. + - Декораторы меняются с @app.* на @router.*. + - Для is_alive добавить prefix="/is_alive" и tags=["healthcheck"]. + - Для post_ws добавить prefix="/ws". + +ЗАДАЧА 2.2 — routes/post_by_file_FORM.py, routes/post_by_url.py + - Аналогичный перевод на APIRouter. + - Добавить tags=["ASR"]. + - Убрать импорт app из utils.pre_start_init. + +ЗАДАЧА 2.3 — routes/ws_audio_transkrib.py, routes/demo_page.py + - Перевод WebSocket и HTML-роута на APIRouter. + - Очистка неиспользуемых импортов (subprocess, WebSocketException). + - Убрать app.mount("/static", ...) из demo_page.py (перенесено в main.py). + +После завершения этого этапа в main.py появляется единственное место +подключения всех роутеров через app.include_router(...). + +================================================================================ +ЭТАП 3. МОДЕЛИ API: ЕДИНЫЙ КОНТРАКТ ОТВЕТОВ +================================================================================ + +Цель: стандартизировать формат ответов API и заложить структуру для будущей +авторизации. + +ЗАДАЧА 3.1 — models/fast_api_models.py + - Добавить базовые Pydantic-модели: + * BaseResponse / ErrorResponse (унификация полей success, + error_description, data). + * UserBase, Token, TokenPayload — заготовки для JWT-аутентификации. + - Реорганизовать существующие SyncASRRequest, PostFileRequest, + PostFileRequestDiarize, WebSocketModel без потери функциональности. + - Добавить примеры ответов (json_schema_extra) для документации. + +================================================================================ +ЭТАП 4. ИНФРАСТРУКТУРА HIGHLOAD (ПОДГОТОВКА) +================================================================================ + +Цель: подготовить приложение к работе за reverse-proxy, в k8s и под нагрузкой. + +ЗАДАЧА 4.1 — main.py (дополнение) + - Добавить middleware: + * TrustedHostMiddleware. + * Обработчик ошибок (HTTPException -> JSON-ответ). + * Генерация request_id для трейсинга (или через middleware, + или через correlation_id). + - Настроить корректные заголовки для проксирования (X-Forwarded-For). + +ЗАДАЧА 4.2 — routes/is_alive.py (расширение) + - Разделить эндпоинт на: + * /health/live (liveness probe) — проверка, что процесс жив. + * /health/ready (readiness probe) — проверка готовности принимать + трафик (модели загружены, память GPU в норме). + - Добавить метрики: uptime, количество обработанных запросов (заготовка). + +ЗАДАЧА 4.3 — создание core/logging_config.py (новый файл) + - Заготовка структурированного JSON-логирования. + - Формат: {"timestamp": "...", "level": "...", "message": "...", + "request_id": "...", "module": "..."}. + - Подготовка для последующей интеграции с ELK / Loki / Grafana. + +================================================================================ +ЭТАП 5. ПОДГОТОВКА К AUTH, ADMIN, PAYMENTS +================================================================================ + +Цель: создать слой security и структуру под будущие модули, не ломая текущий +ASR-функционал. + +ЗАДАЧА 5.1 — создание core/security.py (новый файл) + - Утилиты для хеширования паролей (passlib / bcrypt). + - Создание и верификация JWT-токенов (jose / PyJWT). + - SECRET_KEY, ALGORITHM, ACCESS_TOKEN_EXPIRE_MINUTES из config. + +ЗАДАЧА 5.2 — создание api/deps.py (новый файл) + - Зависимости FastAPI (Depends): + * get_current_user — извлечение пользователя из токена. + * get_current_active_user — проверка статуса пользователя. + - Пока реализовать как заглушки (stub), возвращающие фиктивного пользователя. + +ЗАДАЧА 5.3 — создание структуры api/v1/ (новые директории) + - Перенос текущих роутов под префикс /api/v1/. + - Освобождение корневых путей (/, /demo, /docs) для админки и фронта. + - Создание api/v1/endpoints/ для ASR-роутов. + - Создание api/v1/api.py для агрегации роутеров v1. + +ЗАДАЧА 5.4 — models/fast_api_models.py (дополнение) + - Добавить модели для платежной системы (заготовки): + * Subscription (тип подписки, дата начала/окончания, статус). + * Transaction (id, сумма, валюта, статус, внешний id платежа). + * UserAccount (баланс, тариф). + - Добавить enum-ы для статусов. + +================================================================================ + ПОРЯДОК СТАРТА +================================================================================ + +Стартовать рекомендуется с Этапа 1, Задача 1.1 (config.py). +После согласования каждой задачи выдается полный текст измененного файла +(или согласованной группы файлов) для замены в репозитории. + +================================================================================ From 2e7e82157bf69235d430765c9318d62ac05d7d45 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 12:32:18 +0300 Subject: [PATCH 02/46] =?UTF-8?q?refactor:=20=D0=BF=D0=B5=D1=80=D0=B5?= =?UTF-8?q?=D0=B2=D0=B5=D1=81=D1=82=D0=B8=20config.py=20=D0=BD=D0=B0=20pyd?= =?UTF-8?q?antic-settings?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: aider (ollama/kimi-k2.6:latest-cloud) --- config.py | 230 +++++++++++++++++++++++++++++++++++------------------- 1 file changed, 151 insertions(+), 79 deletions(-) diff --git a/config.py b/config.py index cf19b32..6f4f6be 100644 --- a/config.py +++ b/config.py @@ -1,80 +1,155 @@ import os -import datetime -from tarfile import DEFAULT_FORMAT - -# server settings -HOST = os.getenv('HOST', '0.0.0.0') -PORT = int(os.getenv('PORT', 49153)) - -# Model settings -MODEL_NAME = os.getenv('MODEL_NAME', "gigaam-v3-ctc") ## Vosk5SmallStreaming Vosk5 Gigaam Whisper Gigaam_rnnt, "gigaam-v3-rnnt", "gigaam-v3-ctc" -BASE_SAMPLE_RATE = int(os.getenv('BASE_SAMPLE_RATE', 16000)) # Стрим из астериска отдаёт только 8к -PROVIDER = os.getenv('PROVIDER',"CUDA") -NUM_THREADS = int(os.getenv('NUM_THREADS', 0)) - -# HuggingFaceHubSettings -os.environ["HF_HOME"] = os.getenv("HF_HOME", "./models") - -# Logger settings -LOGGING_LEVEL = os.getenv('LOGGING_LEVEL', 'DEBUG') -LOGGING_FORMAT = os.getenv('LOGGING_FORMAT', u'#%(levelname)-8s %(filename)s [LINE:%(lineno)d] [%(asctime)s] %(message)s') -FILENAME = os.getenv('FILENAME', f'logs/ASR-{datetime.datetime.now().date()}.log') -FILEMODE = os.getenv('FILEMODE', 'a') -LOG_BACKUP_COUNT = os.getenv('LOG_BACKUP_COUNT', 180) # Срок хранения логов в днях -IS_PROD = True if int(os.getenv('IS_PROD', 1))==1 else False - -# Recognition_settings -MAX_OVERLAP_DURATION = int(os.getenv('MAX_OVERLAP_DURATION', 30)) # Максимальная продолжительность буфера аудио (зависит от модели) приемлемый диапазон 10-15 сек. Для Vosk, для Гига СТС можно больше. -RECOGNITION_ATTEMPTS = 1 # Пока не менять -SPEECH_PER_SEC_NORM_RATE = 18 # Нормальное количество токенов в секунду. При превышении этого значения становится -# возможным автоматически замедлять скорость речи для улучшения распознавания. В реальной речи, как правило, находится -# в интервале от 13 до 25. -MAKE_MONO = True if int(os.getenv('MAKE_MONO', 0)) == 1 else False -USE_BATCH = True if int(os.getenv('USE_BATCH', 1)) == 1 else False -ASR_BATCH_SIZE = int(os.getenv('ASR_BATCH_SIZE', 8)) # Размер батча для распознавания аудио. - -# Vad_settings -VAD_SENSITIVITY = int(os.getenv('VAD_SENSE', 3)) # 1 to 5 Higher - more words. -VAD_WITH_GPU = True if int(os.getenv('VAD_WITH_GPU', 0)) == 1 else False - -# Sentensize_settings -BETWEEN_WORDS_PERCENTILE = int(os.getenv('BETWEEN_WORDS_PERCENTILE', 80)) # Параметр определяет как мелко будет биться -# текст на предложения. Чем меньше значение, тем более короткие будут предложения. В среднем в одном предложении 10 слов. -# То есть, по длительности каждая десятая пауза означает конец предложения или мысли. Влияет на пунктуацию выражений. - -# Punctuate_settings -CAN_PUNCTUATE = True if int(os.getenv('CAN_PUNCTUATE', 1)) == 1 else False -PUNCTUATE_WITH_GPU = True if int(os.getenv('PUNCTUATE_WITH_GPU', 0)) == 1 else False - -# Diarisation_settings -CAN_DIAR = True if int(os.getenv('CAN_DIAR', 0)) == 1 else False -DIAR_MODEL_NAME = str(os.getenv('DIAR_MODEL_NAME', "voxblink2_samresnet100_ft")+".onnx") -DIAR_WITH_GPU = True if (int(os.getenv('DIAR_WITH_GPU', 0)) == 1 and PROVIDER in ["CUDA", "TENSORRT"]) else False -CPU_WORKERS = int(os.getenv('CPU_WORKERS', 0)) # Для значений меньше 1 будут использованы все доступные ядра. -# При значении от 1 - указанное число ядер CPU. Работает только при DIAR_WITH_GPU False -DIAR_GPU_BATCH_SIZE = int(os.getenv('DIAR_GPU_BATCH_SIZE', 2)) # Ширина Батча для процесса извлечения эмбеддингов с GPU. -# Оптимально от 4 до 16. Дальнейшее увеличение приводит к неоправданному расходу памяти. - -# Разных моделей для диаризации много. -# [('cnceleb_resnet34', 25), ('cnceleb_resnet34_LM', 25), ('voxblink2_samresnet100', 191), ('voxblink2_samresnet100_ft', 191), -# ('voxblink2_samresnet34', 96), ('voxblink2_samresnet34_ft', 96), ('voxceleb_CAM++', 27), ('voxceleb_CAM++_LM', 27), -# ('voxceleb_ECAPA1024', 56), ('voxceleb_ECAPA1024_LM', 56), ('voxceleb_ECAPA512', 23), ('voxceleb_ECAPA512_LM', 23), -# ('voxceleb_gemini_dfresnet114_LM', 24), ('voxceleb_resnet152_LM', 75), ('voxceleb_resnet221_LM', 90), -# ('voxceleb_resnet293_LM', 109), ('voxceleb_resnet34', 25), ('voxceleb_resnet34_LM', 25)] - -# Инструменты управления распознаванием быстрой речи. -DO_SPEED_SPEECH_CORRECTION = True if int(os.getenv('USE_SPEED_SPEECH_CORRECTION', 1)) == 1 else False # Включено - -# 1 - обычная скорость, меньше - медленнее, больше - быстрее -SPEED_SPEECH_CORRECTION_MULTIPLIER = float(os.getenv('SPEED_SPEECH_CORRECTION_MULTIPLIER', 1)) - -# Настройки сервиса локального распознавания. -DO_LOCAL_FILE_RECOGNITIONS = True if int(os.getenv('DO_LOCAL_FILE_RECOGNITIONS', 0)) == 1 else False -DELETE_LOCAL_FILE_AFTR_ASR = True if int(os.getenv('DELETE_LOCAL_FILE_AFTR_ASR', 0)) == 1 else False -HUMAN_FORMAT_MD_FILE = True if int(os.getenv('HUMAN_FORMAT_MD_FILE', 0)) == 1 else False - - -AUDIOEXTENTIONS = [ +from datetime import date +from pydantic_settings import BaseSettings, SettingsConfigDict +from pydantic import Field, field_validator, model_validator + + +class Settings(BaseSettings): + """ + Конфигурация приложения на базе pydantic-settings. + Все значения могут быть переопределены через переменные окружения + или файл .env в корне проекта. + """ + model_config = SettingsConfigDict( + env_file='.env', + env_file_encoding='utf-8', + extra='ignore', + ) + + # server settings + HOST: str = '0.0.0.0' + PORT: int = 49153 + + # Model settings + MODEL_NAME: str = "gigaam-v3-ctc" + BASE_SAMPLE_RATE: int = 16000 + PROVIDER: str = "CUDA" + NUM_THREADS: int = 0 + + # HuggingFace Hub settings + HF_HOME: str = "./models" + + # Logger settings + LOGGING_LEVEL: str = 'DEBUG' + LOGGING_FORMAT: str = '#%(levelname)-8s %(filename)s [LINE:%(lineno)d] [%(asctime)s] %(message)s' + FILENAME: str = Field(default_factory=lambda: f'logs/ASR-{date.today()}.log') + FILEMODE: str = 'a' + LOG_BACKUP_COUNT: int = 180 + IS_PROD: bool = True + + # Recognition settings + MAX_OVERLAP_DURATION: int = 30 + RECOGNITION_ATTEMPTS: int = 1 + SPEECH_PER_SEC_NORM_RATE: int = 18 + MAKE_MONO: bool = False + USE_BATCH: bool = True + ASR_BATCH_SIZE: int = 8 + + # VAD settings + VAD_SENSITIVITY: int = 3 + VAD_WITH_GPU: bool = False + + # Sentensize settings + BETWEEN_WORDS_PERCENTILE: int = 80 + + # Punctuate settings + CAN_PUNCTUATE: bool = True + PUNCTUATE_WITH_GPU: bool = False + + # Diarisation settings + CAN_DIAR: bool = False + DIAR_MODEL_NAME: str = "voxblink2_samresnet100_ft" + DIAR_WITH_GPU: bool = False + CPU_WORKERS: int = 0 + DIAR_GPU_BATCH_SIZE: int = 2 + + # Speed speech correction + DO_SPEED_SPEECH_CORRECTION: bool = True + SPEED_SPEECH_CORRECTION_MULTIPLIER: float = 1.0 + + # Local file recognition + DO_LOCAL_FILE_RECOGNITIONS: bool = False + DELETE_LOCAL_FILE_AFTR_ASR: bool = False + HUMAN_FORMAT_MD_FILE: bool = False + + @field_validator( + 'IS_PROD', 'MAKE_MONO', 'USE_BATCH', 'VAD_WITH_GPU', + 'CAN_PUNCTUATE', 'PUNCTUATE_WITH_GPU', 'CAN_DIAR', + 'DIAR_WITH_GPU', 'DO_SPEED_SPEECH_CORRECTION', + 'DO_LOCAL_FILE_RECOGNITIONS', 'DELETE_LOCAL_FILE_AFTR_ASR', + 'HUMAN_FORMAT_MD_FILE', + mode='before' + ) + @classmethod + def _int_to_bool(cls, v): + if isinstance(v, bool): + return v + if isinstance(v, int): + return v == 1 + if isinstance(v, str): + return int(v) == 1 + return bool(v) + + @model_validator(mode='after') + def _compute_derived(self): + # Сохраняем side-effect для HuggingFace (совместимость с существующим кодом) + os.environ["HF_HOME"] = self.HF_HOME + + # К имени модели диаризации всегда добавляем расширение .onnx + if not self.DIAR_MODEL_NAME.endswith('.onnx'): + self.DIAR_MODEL_NAME = self.DIAR_MODEL_NAME + '.onnx' + + # DIAR_WITH_GPU актуален только при GPU-провайдерах + if self.DIAR_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.DIAR_WITH_GPU = False + + return self + + +# Единственный экземпляр настроек +settings = Settings() + +# ============================================================================== +# Обратная совместимость: экспортируем атрибуты на уровень модуля. +# Все остальные модули проекта могут продолжать использовать `config.HOST`, +# `config.PORT` и т.д. без изменений. +# ============================================================================== +HOST = settings.HOST +PORT = settings.PORT +MODEL_NAME = settings.MODEL_NAME +BASE_SAMPLE_RATE = settings.BASE_SAMPLE_RATE +PROVIDER = settings.PROVIDER +NUM_THREADS = settings.NUM_THREADS +HF_HOME = settings.HF_HOME +LOGGING_LEVEL = settings.LOGGING_LEVEL +LOGGING_FORMAT = settings.LOGGING_FORMAT +FILENAME = settings.FILENAME +FILEMODE = settings.FILEMODE +LOG_BACKUP_COUNT = settings.LOG_BACKUP_COUNT +IS_PROD = settings.IS_PROD +MAX_OVERLAP_DURATION = settings.MAX_OVERLAP_DURATION +RECOGNITION_ATTEMPTS = settings.RECOGNITION_ATTEMPTS +SPEECH_PER_SEC_NORM_RATE = settings.SPEECH_PER_SEC_NORM_RATE +MAKE_MONO = settings.MAKE_MONO +USE_BATCH = settings.USE_BATCH +ASR_BATCH_SIZE = settings.ASR_BATCH_SIZE +VAD_SENSITIVITY = settings.VAD_SENSITIVITY +VAD_WITH_GPU = settings.VAD_WITH_GPU +BETWEEN_WORDS_PERCENTILE = settings.BETWEEN_WORDS_PERCENTILE +CAN_PUNCTUATE = settings.CAN_PUNCTUATE +PUNCTUATE_WITH_GPU = settings.PUNCTUATE_WITH_GPU +CAN_DIAR = settings.CAN_DIAR +DIAR_MODEL_NAME = settings.DIAR_MODEL_NAME +DIAR_WITH_GPU = settings.DIAR_WITH_GPU +CPU_WORKERS = settings.CPU_WORKERS +DIAR_GPU_BATCH_SIZE = settings.DIAR_GPU_BATCH_SIZE +DO_SPEED_SPEECH_CORRECTION = settings.DO_SPEED_SPEECH_CORRECTION +SPEED_SPEECH_CORRECTION_MULTIPLIER = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER +DO_LOCAL_FILE_RECOGNITIONS = settings.DO_LOCAL_FILE_RECOGNITIONS +DELETE_LOCAL_FILE_AFTR_ASR = settings.DELETE_LOCAL_FILE_AFTR_ASR +HUMAN_FORMAT_MD_FILE = settings.HUMAN_FORMAT_MD_FILE + +AUDIOEXTENTIONS = [ # Основные форматы '.mp3', '.wav', '.aac', '.ogg', '.flac', '.m4a', '.wma', '.aiff', '.alac', # Менее распространённые форматы @@ -88,6 +163,3 @@ # Редкие/устаревшие форматы '.669', '.mtm', '.med', '.far', '.umx' ] - - -print(f"Using '{LOGGING_LEVEL}' LOGGING_LEVEL") \ No newline at end of file From c4351ca2de6f313dd29469a1fdfdca94ccd63ae0 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 12:52:21 +0300 Subject: [PATCH 03/46] =?UTF-8?q?fix:=20=D0=B4=D0=BE=D0=B1=D0=B0=D0=B2?= =?UTF-8?q?=D0=B8=D1=82=D1=8C=20=D0=B2=D0=B0=D0=BB=D0=B8=D0=B4=D0=B0=D1=86?= =?UTF-8?q?=D0=B8=D1=8E=20GPU-=D1=84=D0=BB=D0=B0=D0=B3=D0=BE=D0=B2=20VAD?= =?UTF-8?q?=20=D0=B8=20=D0=BF=D1=83=D0=BD=D0=BA=D1=82=D1=83=D0=B0=D1=86?= =?UTF-8?q?=D0=B8=D0=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/config.py b/config.py index 6f4f6be..d15379a 100644 --- a/config.py +++ b/config.py @@ -103,6 +103,15 @@ def _compute_derived(self): if self.DIAR_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: self.DIAR_WITH_GPU = False + # VAD_WITH_GPU актуален только при GPU-провайдерах + if self.VAD_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.VAD_WITH_GPU = False + + # DIAR_WITH_GPU актуален только при GPU-провайдерах + if self.PUNCTUATE_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.PUNCTUATE_WITH_GPU = False + + return self From f927a1cf3d4ff8c432bdd605d2f0b60ad638dd63 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 13:22:54 +0300 Subject: [PATCH 04/46] fix: correct audio config permissions --- config.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/config.py b/config.py index d15379a..e8f2a72 100644 --- a/config.py +++ b/config.py @@ -171,4 +171,4 @@ def _compute_derived(self): '.xm', '.mod', '.s3m', '.it', '.nsf', # Редкие/устаревшие форматы '.669', '.mtm', '.med', '.far', '.umx' -] +] \ No newline at end of file From c5f85fb3eddf2860071e27636b40aa73857d3ba8 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 13:24:39 +0300 Subject: [PATCH 05/46] =?UTF-8?q?docs:=20=D0=BF=D0=B5=D1=80=D0=B5=D0=BD?= =?UTF-8?q?=D0=B5=D1=81=D1=82=D0=B8=20=D0=BA=D0=BE=D0=BC=D0=BC=D0=B5=D0=BD?= =?UTF-8?q?=D1=82=D0=B0=D1=80=D0=B8=D0=B8=20=D0=BA=20=D0=BF=D0=B5=D1=80?= =?UTF-8?q?=D0=B5=D0=BC=D0=B5=D0=BD=D0=BD=D1=8B=D0=BC=20=D0=B2=20config.py?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: aider (ollama/kimi-k2.6:latest-cloud) --- config.py | 35 +++++++++++++++++++++++++++++++---- 1 file changed, 31 insertions(+), 4 deletions(-) diff --git a/config.py b/config.py index e8f2a72..d8f89a7 100644 --- a/config.py +++ b/config.py @@ -21,7 +21,9 @@ class Settings(BaseSettings): PORT: int = 49153 # Model settings + # Vosk5SmallStreaming Vosk5 Gigaam Whisper Gigaam_rnnt, "gigaam-v3-rnnt", "gigaam-v3-ctc" MODEL_NAME: str = "gigaam-v3-ctc" + # Стрим из астериска отдаёт только 8к BASE_SAMPLE_RATE: int = 16000 PROVIDER: str = "CUDA" NUM_THREADS: int = 0 @@ -34,22 +36,35 @@ class Settings(BaseSettings): LOGGING_FORMAT: str = '#%(levelname)-8s %(filename)s [LINE:%(lineno)d] [%(asctime)s] %(message)s' FILENAME: str = Field(default_factory=lambda: f'logs/ASR-{date.today()}.log') FILEMODE: str = 'a' + # Срок хранения логов в днях LOG_BACKUP_COUNT: int = 180 IS_PROD: bool = True # Recognition settings + # Максимальная продолжительность буфера аудио (зависит от модели) приемлемый диапазон 10-15 сек. + # Для Vosk, для Гига СТС можно больше. MAX_OVERLAP_DURATION: int = 30 + # Пока не менять RECOGNITION_ATTEMPTS: int = 1 + # Нормальное количество токенов в секунду. При превышении этого значения становится + # возможным автоматически замедлять скорость речи для улучшения распознавания. В реальной речи, как правило, находится + # в интервале от 13 до 25. SPEECH_PER_SEC_NORM_RATE: int = 18 MAKE_MONO: bool = False USE_BATCH: bool = True + # Размер батча для распознавания аудио. ASR_BATCH_SIZE: int = 8 - # VAD settings + # Vad settings + # 1 to 5 Higher - more words. VAD_SENSITIVITY: int = 3 VAD_WITH_GPU: bool = False # Sentensize settings + # Параметр определяет как мелко будет биться текст на предложения. Чем меньше значение, + # тем более короткие будут предложения. В среднем в одном предложении 10 слов. + # То есть, по длительности каждая десятая пауза означает конец предложения или мысли. + # Влияет на пунктуацию выражений. BETWEEN_WORDS_PERCENTILE: int = 80 # Punctuate settings @@ -58,16 +73,29 @@ class Settings(BaseSettings): # Diarisation settings CAN_DIAR: bool = False + # Разных моделей для диаризации много. + # [('cnceleb_resnet34', 25), ('cnceleb_resnet34_LM', 25), ('voxblink2_samresnet100', 191), ('voxblink2_samresnet100_ft', 191), + # ('voxblink2_samresnet34', 96), ('voxblink2_samresnet34_ft', 96), ('voxceleb_CAM++', 27), ('voxceleb_CAM++_LM', 27), + # ('voxceleb_ECAPA1024', 56), ('voxceleb_ECAPA1024_LM', 56), ('voxceleb_ECAPA512', 23), ('voxceleb_ECAPA512_LM', 23), + # ('voxceleb_gemini_dfresnet114_LM', 24), ('voxceleb_resnet152_LM', 75), ('voxceleb_resnet221_LM', 90), + # ('voxceleb_resnet293_LM', 109), ('voxceleb_resnet34', 25), ('voxceleb_resnet34_LM', 25)] DIAR_MODEL_NAME: str = "voxblink2_samresnet100_ft" DIAR_WITH_GPU: bool = False + # Для значений меньше 1 будут использованы все доступные ядра. + # При значении от 1 - указанное число ядер CPU. Работает только при DIAR_WITH_GPU False CPU_WORKERS: int = 0 + # Ширина Батча для процесса извлечения эмбеддингов с GPU. + # Оптимально от 4 до 16. Дальнейшее увеличение приводит к неоправданному расходу памяти. DIAR_GPU_BATCH_SIZE: int = 2 # Speed speech correction + # Инструменты управления распознаванием быстрой речи. Включено DO_SPEED_SPEECH_CORRECTION: bool = True + # 1 - обычная скорость, меньше - медленнее, больше - быстрее SPEED_SPEECH_CORRECTION_MULTIPLIER: float = 1.0 # Local file recognition + # Настройки сервиса локального распознавания. DO_LOCAL_FILE_RECOGNITIONS: bool = False DELETE_LOCAL_FILE_AFTR_ASR: bool = False HUMAN_FORMAT_MD_FILE: bool = False @@ -107,11 +135,10 @@ def _compute_derived(self): if self.VAD_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: self.VAD_WITH_GPU = False - # DIAR_WITH_GPU актуален только при GPU-провайдерах + # PUNCTUATE_WITH_GPU актуален только при GPU-провайдерах if self.PUNCTUATE_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: self.PUNCTUATE_WITH_GPU = False - return self @@ -171,4 +198,4 @@ def _compute_derived(self): '.xm', '.mod', '.s3m', '.it', '.nsf', # Редкие/устаревшие форматы '.669', '.mtm', '.med', '.far', '.umx' -] \ No newline at end of file +] From 0a47d690ccadfe1f6f24ff6a656a7f29795514c6 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 16:55:16 +0300 Subject: [PATCH 06/46] =?UTF-8?q?=D0=9F=D1=80=D0=B0=D0=B2=D0=BA=D0=B0=20AP?= =?UTF-8?q?P=20=D0=B8=20=D1=86=D0=B5=D0=BD=D1=82=D1=80=D0=B0=D0=BB=D0=B8?= =?UTF-8?q?=D0=B7=D0=B0=D1=86=D0=B8=D1=8F=20=D0=B2=20main.py.=20=D0=9F?= =?UTF-8?q?=D0=B5=D1=80=D0=B5=D0=B5=D0=B7=D0=B4=20=D1=80=D0=BE=D1=83=D1=82?= =?UTF-8?q?=D0=B5=D1=80=D0=BE=D0=B2=20=D0=BD=D0=B0=20APIrouter?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 3 ++ global_to_do.txt | 6 ++-- main.py | 62 ++++++++++++++++++++++++++++++++++-- requirements.txt | 3 +- routes/demo_page.py | 9 ++---- routes/is_alive.py | 8 +++-- routes/post_by_file_FORM.py | 8 +++-- routes/post_by_url.py | 9 ++++-- routes/post_ws.py | 6 ++-- routes/root.py | 8 +++-- routes/ws_audio_transkrib.py | 9 +++--- utils/pre_start_init.py | 28 ---------------- 12 files changed, 99 insertions(+), 60 deletions(-) diff --git a/.gitignore b/.gitignore index 6d883c2..1a208b8 100644 --- a/.gitignore +++ b/.gitignore @@ -17,3 +17,6 @@ /models/sherpa-onnx-whisper-small/ /trash/test_data.py /.idea/ +.aider* +/.continue/ +.env diff --git a/global_to_do.txt b/global_to_do.txt index 6cd4622..61fdeeb 100644 --- a/global_to_do.txt +++ b/global_to_do.txt @@ -17,14 +17,14 @@ Цель: устранить side-effects при импорте, валидировать окружение и централизовать создание приложения. -ЗАДАЧА 1.1 — config.py +ЗАДАЧА 1.1 — config.py [ВЫПОЛНЕНО] - Миграция с голого os.getenv() на pydantic-settings (BaseSettings). - Все переменные получают типизацию, значения по умолчанию и валидацию. - Старые имена экспортируются для обратной совместимости с неизменяемыми модулями ASR (Recognizer, Diarisation и т.д.). - Добавить SettingsConfigDict с поддержкой .env файла. -ЗАДАЧА 1.2 — utils/pre_start_init.py +ЗАДАЧА 1.2 — utils/pre_start_init.py [ВЫПОЛНЕНО] - Удалить глобальный объект app = FastAPI(...). - Оставить только инициализацию paths, глобальных defaultdict (audio_overlap, audio_buffer, audio_to_asr, audio_duration, @@ -33,7 +33,7 @@ - Убрать любие импорты роутов или HTTP-объектов. - Файл должен быть безопасен для импорта без побочных эффектов. -ЗАДАЧА 1.3 — main.py +ЗАДАЧА 1.3 — main.py [ВЫПОЛНЕНО] - Создание app = FastAPI(...) переносится сюда как явная фабрика. - Подключение lifespan из utils.pre_start_init. - Добавление CORS-middleware (CORSMiddleware). diff --git a/main.py b/main.py index b1b5695..5acd0cb 100644 --- a/main.py +++ b/main.py @@ -1,10 +1,68 @@ from utils.do_logging import logger import uvicorn import config -from utils.pre_start_init import app -import routes, models +from contextlib import asynccontextmanager +from fastapi import FastAPI +from fastapi.middleware.cors import CORSMiddleware +from fastapi.staticfiles import StaticFiles from fastapi.openapi.utils import get_openapi +from utils.files_whatcher import start_file_watcher +from utils.pre_start_init import paths +import threading +from routes.root import router as root_router +from routes.is_alive import router as is_alive_router +from routes.post_ws import router as post_ws_router +from routes.post_by_file_FORM import router as post_by_file_router +from routes.post_by_url import router as post_by_url_router +from routes.ws_audio_transkrib import router as ws_audio_transkrib_router +from routes.demo_page import router as demo_router +import models + + +@asynccontextmanager +async def lifespan(app): + # on_start + logger.debug("Приложение FastAPI запущено") + if config.DO_LOCAL_FILE_RECOGNITIONS: + observer_thread = threading.Thread( + target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), + daemon=True + ) + observer_thread.start() + logger.info("File watcher started") + + yield # Здесь приложение работает + + +app = FastAPI( + lifespan=lifespan, + version="1.0", + docs_url='/docs', + root_path='/root', + title='ASR on SHERPA-ONNX' +) + +# CORS middleware +app.add_middleware( + CORSMiddleware, + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + +# Static files +app.mount("/static", StaticFiles(directory="static"), name="static") + +# Routers +app.include_router(root_router) +app.include_router(is_alive_router) +app.include_router(post_ws_router) +app.include_router(post_by_file_router) +app.include_router(post_by_url_router) +app.include_router(ws_audio_transkrib_router) +app.include_router(demo_router) def custom_openapi(): openapi_schema = get_openapi( diff --git a/requirements.txt b/requirements.txt index c00d457..ed26c38 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,5 @@ setuptools +pydantic_settings httpx~=0.28.1 pathlib~=1.0.1 @@ -20,8 +21,6 @@ tqdm~=4.67.1 psutil~=7.0.0 - - librosa~=0.11.0 # For punctuation diff --git a/routes/demo_page.py b/routes/demo_page.py index fa01f14..7b13da8 100644 --- a/routes/demo_page.py +++ b/routes/demo_page.py @@ -1,19 +1,16 @@ -from utils.pre_start_init import app -from fastapi import WebSocket, WebSocketException, Request +from fastapi import APIRouter, Request from utils.do_logging import logger -from fastapi.staticfiles import StaticFiles from fastapi.responses import HTMLResponse from fastapi.templating import Jinja2Templates -# Mount static files -app.mount("/static", StaticFiles(directory="static"), name="static") +router = APIRouter() # Setup templates templates = Jinja2Templates(directory="templates") -@app.get("/demo", response_class=HTMLResponse) +@router.get("/demo", response_class=HTMLResponse) async def demo_page(request: Request): return templates.TemplateResponse( "index.html", diff --git a/routes/is_alive.py b/routes/is_alive.py index bdd2aa4..d1735e2 100644 --- a/routes/is_alive.py +++ b/routes/is_alive.py @@ -1,4 +1,4 @@ -from utils.pre_start_init import app +from fastapi import APIRouter import logging import datetime import os @@ -6,6 +6,8 @@ from utils.pre_start_init import audio_to_asr +router = APIRouter() + def get_gpu_free_memory(): try: pynvml.nvmlInit() @@ -22,7 +24,7 @@ def get_gpu_free_memory(): return free_mb, gpu_load,temperature -@app.get("/is_alive") +@router.get("/is_alive") async def check_if_service_is_alive(): logging.info('GET_is_alive') @@ -42,4 +44,4 @@ async def check_if_service_is_alive(): "free_memory_mb": free_mb, "gpu_load_percent": gpu_load, "temperature_celsius": temperature - } \ No newline at end of file + } diff --git a/routes/post_by_file_FORM.py b/routes/post_by_file_FORM.py index 5d8dce9..4e9aba8 100644 --- a/routes/post_by_file_FORM.py +++ b/routes/post_by_file_FORM.py @@ -2,14 +2,16 @@ import asyncio import config -from utils.pre_start_init import app +from fastapi import APIRouter, Depends, File, Form, UploadFile from utils.do_logging import logger from models.fast_api_models import PostFileRequest from Recognizer.engine.file_recognition import process_file -from fastapi import Depends, File, Form, UploadFile from threading import Lock +router = APIRouter() + + # Глобальный лок для потокобезопасности audio_lock = Lock() @@ -42,7 +44,7 @@ def get_file_request( ) -@app.post("/post_file") +@router.post("/post_file") async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), diff --git a/routes/post_by_url.py b/routes/post_by_url.py index e05e516..5305658 100644 --- a/routes/post_by_url.py +++ b/routes/post_by_url.py @@ -1,7 +1,8 @@ import uuid import asyncio import os -from utils.pre_start_init import app, posted_and_downloaded_audio +from fastapi import APIRouter +from utils.pre_start_init import posted_and_downloaded_audio from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest @@ -10,10 +11,12 @@ from io import BytesIO +router = APIRouter() + # Глобальный лок для потокобезопасности audio_lock = Lock() -@app.post("/post_one_step_req") +@router.post("/post_one_step_req") async def post(params: SyncASRRequest): """ На вход принимает HttpUrl - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n @@ -60,4 +63,4 @@ async def post(params: SyncASRRequest): result["success"] = False result['error_description'] = str(error_description) - return result \ No newline at end of file + return result diff --git a/routes/post_ws.py b/routes/post_ws.py index 6c674f4..be09dfd 100644 --- a/routes/post_ws.py +++ b/routes/post_ws.py @@ -1,7 +1,9 @@ -from utils.pre_start_init import app +from fastapi import APIRouter from models.fast_api_models import WebSocketModel -@app.post("/ws") +router = APIRouter() + +@router.post("/ws") async def post_not_websocket(ws:WebSocketModel): """Описание для вебсокета ниже в описании WebSocketModel """ return f"Прочти инструкцию в Schemas - 'WebSocketModel'" diff --git a/routes/root.py b/routes/root.py index f8749b9..73f7cfd 100644 --- a/routes/root.py +++ b/routes/root.py @@ -1,11 +1,13 @@ -from utils.pre_start_init import app +from fastapi import APIRouter import config -@app.get("/") +router = APIRouter() + +@router.get("/") async def root(): print("Зашли в root") return {"error": True, "data": "No_service_selected", # "available_services": ["Vosk_Recognizer"], - "comment": f"try_addr: http://{config.HOST}:{config.PORT}/docs"} \ No newline at end of file + "comment": f"try_addr: http://{config.HOST}:{config.PORT}/docs"} diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index a353802..8b9b221 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -4,11 +4,8 @@ import config import uuid from io import BytesIO -import subprocess - -from utils.pre_start_init import app -from fastapi import WebSocket, WebSocketException +from fastapi import APIRouter, WebSocket from utils.do_logging import logger from utils.chunk_doing import find_last_speech_position from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration,ws_collected_asr_res @@ -20,7 +17,9 @@ from Recognizer.engine.stream_recognition import simple_recognise -@app.websocket("/ws") +router = APIRouter() + +@router.websocket("/ws") async def websocket(ws: WebSocket): wait_null_answers=True client_id = uuid.uuid4() diff --git a/utils/pre_start_init.py b/utils/pre_start_init.py index 447df1f..607acfa 100644 --- a/utils/pre_start_init.py +++ b/utils/pre_start_init.py @@ -1,11 +1,6 @@ # -*- coding: utf-8 -*- -from utils.do_logging import logger -from utils.files_whatcher import start_file_watcher -from contextlib import asynccontextmanager from pathlib import Path -from fastapi import FastAPI import config -import threading import gc @@ -63,26 +58,3 @@ gc.set_threshold(500, # быстрые файлы было 700 5, # средне выживающие файлы было 10 5) # долгожители было 10. - - -@asynccontextmanager -async def lifespan(app: FastAPI): - # on_start - logger.debug("Приложение FastAPI запущено") - global observer - if config.DO_LOCAL_FILE_RECOGNITIONS: - observer_thread = threading.Thread( - target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), - daemon=True - ) - observer_thread.start() - logger.info("File watcher started") - - yield # Здесь приложение работает - -app = FastAPI(lifespan=lifespan, - version="1.0", - docs_url='/docs', - root_path='/root', - title='ASR on SHERPA-ONNX' - ) \ No newline at end of file From 5ae1410f9b0402aa3a508199c826f2ff11ce3edb Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 22 Apr 2026 17:51:20 +0300 Subject: [PATCH 07/46] =?UTF-8?q?=D0=9F=D0=B5=D1=80=D0=B5=D1=80=D0=B0?= =?UTF-8?q?=D0=B1=D0=BE=D1=82=D0=BA=D0=B0=20pydantic=20=D0=BC=D0=BE=D0=B4?= =?UTF-8?q?=D0=B5=D0=BB=D0=B5=D0=B9.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- models/fast_api_models.py | 229 ++++++++++++++++++++++++++------------ 1 file changed, 158 insertions(+), 71 deletions(-) diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 234efa3..2c34455 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -1,10 +1,53 @@ -from pydantic import BaseModel, HttpUrl, Field -from typing import Union, Annotated +from pydantic import BaseModel, HttpUrl, Field, ConfigDict +from typing import Union, Annotated, Optional, Any, List, Dict from fastapi import UploadFile import config +class BaseResponse(BaseModel): + """ + Базовая модель ответа API. + """ + success: bool = True + error_description: Optional[str] = None + data: Optional[Dict[str, Any]] = None + + +class ErrorResponse(BaseModel): + """ + Модель ошибки API. + """ + success: bool = False + error_description: str + data: Optional[Dict[str, Any]] = None + + +class UserBase(BaseModel): + """ + Базовая информация о пользователе (заготовка для JWT). + """ + username: Optional[str] = None + email: Optional[str] = None + is_active: bool = True + + +class Token(BaseModel): + """ + Модель токена доступа. + """ + access_token: str + token_type: str = "bearer" + + +class TokenPayload(BaseModel): + """ + Полезная нагрузка JWT-токена. + """ + sub: Optional[str] = None + exp: Optional[int] = None + + class SyncASRRequest(BaseModel): """ :parameter keep_raw: - Если False, то запрос вернёт только пост-обработанные данные do_punctuation и do_dialogue. @@ -27,7 +70,24 @@ class SyncASRRequest(BaseModel): use_batch: Union[bool, None] = config.USE_BATCH batch_size: Union[int, None] = config.ASR_BATCH_SIZE - + model_config = ConfigDict( + json_schema_extra={ + "example": { + "AudioFileUrl": "https://example.com/audio.wav", + "keep_raw": True, + "do_echo_clearing": True, + "do_dialogue": False, + "do_punctuation": False, + "do_diarization": False, + "make_mono": False, + "diar_vad_sensity": 3, + "do_auto_speech_speed_correction": True, + "speech_speed_correction_multiplier": 1.0, + "use_batch": True, + "batch_size": 8 + } + } + ) class PostFileRequest(BaseModel): @@ -52,75 +112,102 @@ class PostFileRequest(BaseModel): speech_speed_correction_multiplier: Union[float, None] = config.SPEED_SPEECH_CORRECTION_MULTIPLIER make_mono: Union[bool, None] = False -class PostFileRequestDiarize(BaseModel): - """ - Модель для проверки запроса пользователя. - :param keep_raw: Если False, то запрос вернёт только пост-обработанные данные do_punctuation и do_dialogue. - :param do_echo_clearing: Проверяет наличие повторений между каналами. - :param num_speakers: Предполагаемое количество спикеров в разговоре. -1 - значит мы не знаем сколько спикеров и определяем их параметром cluster_threshold. - :param cluster_threshold: Значение от 0 до 1. Чем меньше, тем более чувствительное выделение спикеров (тем их больше) - :param do_punctuation: Расставляет пунктуацию. - """ - keep_raw: bool = True - do_echo_clearing: bool = False - do_punctuation: bool = False - num_speakers: int = -1, - cluster_threshold: float = 0.2 + model_config = ConfigDict( + json_schema_extra={ + "example": { + "keep_raw": True, + "do_echo_clearing": False, + "do_dialogue": False, + "do_punctuation": False, + "do_diarization": False, + "use_batch": True, + "batch_size": 8, + "diar_vad_sensity": 3, + "do_auto_speech_speed_correction": True, + "speech_speed_correction_multiplier": 1.0, + "make_mono": False + } + } + ) class WebSocketModel(BaseModel): """OpenAPI не хочет описывать WS, а я не хочу изучать OPEN API. По этому описание тут. - \n - \n Подключение на порт: 49153 - \n На вход жду поток binary, buffer_size +- 6400, mono, wav. - \n На вход я должен получить словарь {'text': { "config" : { "sample_rate" : any(int/float), "wait_null_answers": Bool, - "do_dialogue": Bool, "do_punctuation": Bool}}} - \n do_punctuation отработает только если do_dialogue = True - \n Далее сообщения с данными {"bytes": binary} - \n По окончании передачи {'text': '{ "eof" : 1}'} - \n Ответ получать в формате: {"silence": Bool,"data": str, "error": None/str, "last_message": Bool, - "sentenced_data": {}} - - - \n Пример ответа "data": { - "result" : [{ - "conf" : 1.000000, - "end" : 3.120000, - "start" : 2.340000, - "word" : "здравствуйте" - }, { - "conf" : 1.000000, - "end" : 3.870000, - "start" : 3.600000, - "word" : "вы" - }, - ... - { - "conf" : 0.994019, - "end" : 11.790000, - "start" : 10.890000, - "word" : "записываются" - }], - "text" : "здравствуйте вы ... записываются" -} - -Пример ответа "sentenced_data": { - "raw_text_sentenced_recognition": "channel_1: Татьяна, добрый день. Меня зовут Ульяна.\nchannel_1: Звоню уточнить по поводу документов.\n", - "list_of_sentenced_recognitions": [ - { - "start": 2.28, - "text": "Татьяна, добрый день. Меня зовут Ульяна.", - "speaker": "channel_1" - }, - { - "start": 8.24, - "text": "Звоню уточнить по поводу документов.", - "speaker": "channel_1" - }, - ], - "full_text_only": [ - "Татьяна, добрый день. Меня зовут Ульяна. Звоню уточнить по поводу документов." - ], - "err_state": null - } + + Подключение на порт: 49153 + На вход жду поток binary, buffer_size +- 6400, mono, wav. + + Протокол обмена сообщениями: + 1. Начальное сообщение (конфигурация): + {'text': { "config" : { "sample_rate" : any(int/float), "wait_null_answers": Bool, + "do_dialogue": Bool, "do_punctuation": Bool}}} + do_punctuation отработает только если do_dialogue = True + + 2. Последующие сообщения с аудио-данными: + {"bytes": binary} + + 3. Последнее сообщение (сигнал окончания передачи): + {'text': '{ "eof" : 1}'} + + Формат ответа от сервера: + {"silence": Bool, "data": str, "error": None/str, "last_message": Bool, + "sentenced_data": {}} + + Пример ответа "data": { + "result" : [{ + "conf" : 1.000000, + "end" : 3.120000, + "start" : 2.340000, + "word" : "здравствуйте" + }, { + "conf" : 1.000000, + "end" : 3.870000, + "start" : 3.600000, + "word" : "вы" + }, + ... + { + "conf" : 0.994019, + "end" : 11.790000, + "start" : 10.890000, + "word" : "записываются" + }], + "text" : "здравствуйте вы ... записываются" + } + + Пример ответа "sentenced_data": { + "raw_text_sentenced_recognition": "channel_1: Татьяна, добрый день. Меня зовут Ульяна.'/n'channel_1: Звоню уточнить по поводу документов.", + "list_of_sentenced_recognitions": [ + { + "start": 2.28, + "text": "Татьяна, добрый день. Меня зовут Ульяна.", + "speaker": "channel_1" + }, + { + "start": 8.24, + "text": "Звоню уточнить по поводу документов.", + "speaker": "channel_1" + }, + ], + "full_text_only": [ + "Татьяна, добрый день. Меня зовут Ульяна. Звоню уточнить по поводу документов." + ], + "err_state": null + } """ - pass \ No newline at end of file + + model_config = ConfigDict( + json_schema_extra={ + "example": { + "text": { + "config": { + "sample_rate": 8000, + "wait_null_answers": False, + "do_dialogue": True, + "do_punctuation": True + } + } + } + } + ) + + pass From e4d316579f3714d85c28a7d57e90c64d125199e1 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 23 Apr 2026 10:27:12 +0300 Subject: [PATCH 08/46] =?UTF-8?q?=D0=9A=D0=BE=D1=81=D0=BC=D0=B5=D1=82?= =?UTF-8?q?=D0=B8=D1=87=D0=B5=D1=81=D0=BA=D0=B8=D0=B5=20=D0=BF=D1=80=D0=B0?= =?UTF-8?q?=D0=B2=D0=BA=D0=B8.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 1 + config.py | 2 +- main.py | 10 ++++++++-- utils/pre_start_init.py | 5 ----- 4 files changed, 10 insertions(+), 8 deletions(-) diff --git a/.gitignore b/.gitignore index 1a208b8..e33dccd 100644 --- a/.gitignore +++ b/.gitignore @@ -20,3 +20,4 @@ .aider* /.continue/ .env +/global_to_do.txt diff --git a/config.py b/config.py index d8f89a7..f5968ac 100644 --- a/config.py +++ b/config.py @@ -120,7 +120,7 @@ def _int_to_bool(cls, v): @model_validator(mode='after') def _compute_derived(self): - # Сохраняем side-effect для HuggingFace (совместимость с существующим кодом) + # HF_HOME должен быть установлен ДО импорта библиотек, использующих HuggingFace Hub os.environ["HF_HOME"] = self.HF_HOME # К имени модели диаризации всегда добавляем расширение .onnx diff --git a/main.py b/main.py index 5acd0cb..f9b5702 100644 --- a/main.py +++ b/main.py @@ -1,6 +1,8 @@ from utils.do_logging import logger import uvicorn import config +import os +import gc from contextlib import asynccontextmanager from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware @@ -24,6 +26,10 @@ async def lifespan(app): # on_start logger.debug("Приложение FastAPI запущено") + + # Настройка сборщика мусора. + gc.set_threshold(500, 5, 5) + if config.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), @@ -40,7 +46,7 @@ async def lifespan(app): version="1.0", docs_url='/docs', root_path='/root', - title='ASR on SHERPA-ONNX' + title='ASR' ) # CORS middleware @@ -70,7 +76,7 @@ def custom_openapi(): version="1.0.0", description="Real-time Russian ASR via WebSocket. Send raw audio chunks (16 kHz, mono).", routes=app.routes, - contact={"email": "your@email.com"}, + contact={"email": "Kojevnikov@amulex.ru"}, ) # Добавляем пример для WebSocket openapi_schema["paths"]["/ws"]["websocket"] = { diff --git a/utils/pre_start_init.py b/utils/pre_start_init.py index 607acfa..e47a069 100644 --- a/utils/pre_start_init.py +++ b/utils/pre_start_init.py @@ -1,7 +1,6 @@ # -*- coding: utf-8 -*- from pathlib import Path import config -import gc BASE_DIR = Path(__file__).resolve().parent.parent @@ -54,7 +53,3 @@ ws_collected_asr_res = defaultdict() posted_and_downloaded_audio = defaultdict() -# Устанавливаем новые пороги сборщика мусора -gc.set_threshold(500, # быстрые файлы было 700 - 5, # средне выживающие файлы было 10 - 5) # долгожители было 10. From ae10ec53383d90b250b9fb310f96958876de16de Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 23 Apr 2026 18:00:17 +0300 Subject: [PATCH 09/46] =?UTF-8?q?=D0=9F=D1=80=D0=B8=D0=B2=D0=B5=D0=B4?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20=D0=BA=D0=BE=D0=BD=D1=84=D0=B8=D0=B3?= =?UTF-8?q?=D0=B0=20=D0=BA=20BaseSettings?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Diarisation/__init__.py | 20 +++---- Diarisation/diarazer.py | 4 +- Punctuation/__init__.py | 4 +- Recognizer/__init__.py | 8 +-- Recognizer/engine/file_recognition.py | 22 ++++---- Recognizer/engine/sentensizer.py | 4 +- Recognizer/engine/stream_recognition.py | 16 +++--- VoiceActivityDetector/__init__.py | 6 +-- config.py | 69 ++++++++++++------------- main.py | 6 +-- models/fast_api_models.py | 18 +++---- routes/is_alive.py | 12 +++-- routes/post_by_file_FORM.py | 6 +-- routes/post_by_url.py | 2 +- routes/root.py | 23 ++++++--- routes/ws_audio_transkrib.py | 14 ++--- utils/bytes_to_samples_audio.py | 5 +- utils/chunk_doing.py | 16 +++--- utils/do_logging.py | 19 +++---- utils/files_whatcher.py | 15 +++--- utils/pre_start_init.py | 5 +- 21 files changed, 152 insertions(+), 142 deletions(-) diff --git a/Diarisation/__init__.py b/Diarisation/__init__.py index f879a5c..e1c4841 100644 --- a/Diarisation/__init__.py +++ b/Diarisation/__init__.py @@ -1,13 +1,13 @@ -import config +from config import settings from utils.pre_start_init import paths from utils.do_logging import logger from VoiceActivityDetector import vad import requests -if config.CAN_DIAR: +if settings.CAN_DIAR: if not paths.get("diar_speaker_model_path").exists(): - logger.error(f"Модель для диаризации не найдена. Предпринимаются попытки скачать {config.DIAR_MODEL_NAME}") + logger.error(f"Модель для диаризации не найдена. Предпринимаются попытки скачать {settings.DIAR_MODEL_NAME}") output_path = paths.get("diar_speaker_model_path") api_url = "https://modelscope.cn/api/v1/datasets/wenet/wespeaker_pretrained_models/oss/tree" @@ -16,11 +16,11 @@ response = requests.get(api_url, headers={"User-Agent": "Mozilla/5.0"}) response.raise_for_status() # 2. Ищем нужный файл - target_file = config.DIAR_MODEL_NAME + target_file = settings.DIAR_MODEL_NAME file_data = next((item for item in response.json()["Data"] if item["Key"] == target_file), None) if not file_data: - logger.error(f"Модели с именем {config.DIAR_MODEL_NAME} в списке возможных для загрузки нет.") + logger.error(f"Модели с именем {settings.DIAR_MODEL_NAME} в списке возможных для загрузки нет.") # Получаем список всех ONNX-моделей onnx_models_with_size = [ (item['Key'].split(".")[0], item['Size'] // (1024 * 1024)) @@ -47,23 +47,23 @@ logger.info(f"Модель успешно загружена : {output_path}") else: - logger.debug(f"Будет использован имеющийся файл {config.DIAR_MODEL_NAME}") + logger.debug(f"Будет использован имеющийся файл {settings.DIAR_MODEL_NAME}") if not paths.get("diar_speaker_model_path").exists(): logger.error(f"Модель для Диаризации отсутствует. Диризация выключена и будет не доступна.") logger.error(f"Скачайте модель со страницы 'https://github.com/wenet-e2e/wespeaker/blob/master/docs/pretrained.md' " f"и поместите в по адресу: {str(paths.get('diar_speaker_model_path'))}") - config.CAN_DIAR = False + settings.CAN_DIAR = False else: from .do_diarize import Diarizer diarizer = Diarizer(embedding_model_path=paths.get("diar_speaker_model_path"), vad=vad, max_phrase_gap=1, - batch_size=config.DIAR_GPU_BATCH_SIZE, - cpu_workers=config.CPU_WORKERS, - use_gpu=config.DIAR_WITH_GPU + batch_size=settings.DIAR_GPU_BATCH_SIZE, + cpu_workers=settings.CPU_WORKERS, + use_gpu=settings.DIAR_WITH_GPU ) logger.info(f"Успешно загружена модель Диаризации") diff --git a/Diarisation/diarazer.py b/Diarisation/diarazer.py index de7f6df..fb1c439 100644 --- a/Diarisation/diarazer.py +++ b/Diarisation/diarazer.py @@ -1,13 +1,13 @@ import datetime -import config +from config import settings from Diarisation.do_diarize import load_and_preprocess_audio from utils.pre_start_init import posted_and_downloaded_audio from utils.do_logging import logger from collections import defaultdict -if config.CAN_DIAR: +if settings.CAN_DIAR: from Diarisation import diarizer async def do_diarizing( diff --git a/Punctuation/__init__.py b/Punctuation/__init__.py index 06c912e..c70548b 100644 --- a/Punctuation/__init__.py +++ b/Punctuation/__init__.py @@ -1,10 +1,10 @@ from utils.pre_start_init import paths from utils.do_logging import logger from .punctuate import SbertPuncCaseOnnx -import config +from config import settings try: - sbertpunc = SbertPuncCaseOnnx(paths.get("punctuation_model_path"), use_gpu = config.PUNCTUATE_WITH_GPU) + sbertpunc = SbertPuncCaseOnnx(paths.get("punctuation_model_path"), use_gpu = settings.PUNCTUATE_WITH_GPU) except Exception as e: logger.error(f"Error getting punctuation model - {e}") else: diff --git a/Recognizer/__init__.py b/Recognizer/__init__.py index 92e3923..1263d0d 100644 --- a/Recognizer/__init__.py +++ b/Recognizer/__init__.py @@ -4,7 +4,7 @@ from utils.do_logging import logger from utils import tokens_to_Result from . import engine -import config +from config import settings import onnxruntime as ort import onnx_asr from onnx_asr.loader import PreprocessorRuntimeConfig, OnnxSessionOptions @@ -18,14 +18,14 @@ class Recognizer: def __init__(self): - self.model_name = config.MODEL_NAME + self.model_name = settings.MODEL_NAME self._post_processor = tokens_to_Result.process_single_token_vocab_output self.preprocessor_providers = list() self.encoding_providers = list() self.resampler_providers = list() self.cpu_preprocessing = False - match config.PROVIDER: + match settings.PROVIDER: case "TENSORRT": self.preprocessor_providers = self.encoding_providers = self.resampler_providers = TENSORRT_providers @@ -104,7 +104,7 @@ def __init__(self): ).with_timestamps() try: - audio = np.random.randn(int(config.MAX_OVERLAP_DURATION * config.BASE_SAMPLE_RATE)).astype(np.float32) + audio = np.random.randn(int(settings.MAX_OVERLAP_DURATION * settings.BASE_SAMPLE_RATE)).astype(np.float32) self._recognizer.recognize([audio]) except Exception as e: logger.error("Ошибка при прогреве модели. Сервис работать не будет. Возможно, модель не поддерживает выбранный провайдер.") diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index 635e315..4dc8848 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -1,7 +1,7 @@ import time from pydub import AudioSegment -import config +from config import settings import asyncio import uuid from utils.pre_start_init import ( @@ -58,7 +58,7 @@ def process_file(tmp_path, params): with audio_lock: if posted_and_downloaded_audio[post_id].duration_seconds < 5: logger.debug(f"На вход передано аудио короче 5 секунд. Будет дополнено тишиной ещё 5 сек.") - posted_and_downloaded_audio[post_id] += AudioSegment.silent(duration=5, frame_rate=config.BASE_SAMPLE_RATE) + posted_and_downloaded_audio[post_id] += AudioSegment.silent(duration=5, frame_rate=settings.BASE_SAMPLE_RATE) except Exception as e: error_description += f"Error len_fixing_file: {e}" logger.error(error_description) @@ -70,10 +70,10 @@ def process_file(tmp_path, params): print(f"Начало проверки фреймрейта {(time.perf_counter()-process_file_start):.4f} сек.") try: with audio_lock: - if posted_and_downloaded_audio[post_id].frame_rate != config.BASE_SAMPLE_RATE: + if posted_and_downloaded_audio[post_id].frame_rate != settings.BASE_SAMPLE_RATE: posted_and_downloaded_audio[post_id] = sync_resample_audiosegment( audio_data=posted_and_downloaded_audio[post_id], - target_sample_rate=config.BASE_SAMPLE_RATE) + target_sample_rate=settings.BASE_SAMPLE_RATE) print(f"Корректировка фреймрейта {(time.perf_counter() - process_file_start):.4f} сек.") except KeyError as e_key: error_description = f"Ошибка обращения по ключу {post_id} при изменения фреймрейта - {e_key}" @@ -95,8 +95,8 @@ def process_file(tmp_path, params): # Подготовительные действия try: with audio_lock: - audio_buffer[post_id] = AudioSegment.silent(1, frame_rate=config.BASE_SAMPLE_RATE) - audio_overlap[post_id] = AudioSegment.silent(1, frame_rate=config.BASE_SAMPLE_RATE) + audio_buffer[post_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + audio_overlap[post_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) audio_duration[post_id] = 0 except Exception as e: error_description = f"Ошибка изменения фреймрейта - {e}" @@ -108,15 +108,15 @@ def process_file(tmp_path, params): result["raw_data"].update({f"channel_{n_channel + 1}": list()}) # Основной процесс перебора чанков для распознавания - overlaps = list(mono_data[::config.MAX_OVERLAP_DURATION * 1000]) # Чанки аудио для распознавания + overlaps = list(mono_data[::settings.MAX_OVERLAP_DURATION * 1000]) # Чанки аудио для распознавания total_chunks = len(overlaps) # Количество чанков, для поиска последнего audio_to_asr[post_id] = list() for idx, overlap in enumerate(overlaps): is_last_chunk = (idx == total_chunks - 1) # Если чанк последний with audio_lock: - if (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) < config.MAX_OVERLAP_DURATION: - silent_secs = config.MAX_OVERLAP_DURATION - (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) - overlap += AudioSegment.silent(silent_secs, frame_rate=config.BASE_SAMPLE_RATE) + if (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) < settings.MAX_OVERLAP_DURATION: + silent_secs = settings.MAX_OVERLAP_DURATION - (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) + overlap += AudioSegment.silent(silent_secs, frame_rate=settings.BASE_SAMPLE_RATE) audio_buffer[post_id] = overlap asyncio.run(find_last_speech_position(post_id, is_last_chunk)) # Последний чанк обрабатывается иначе. @@ -183,7 +183,7 @@ def process_file(tmp_path, params): error_description = f"Error echo clearing - {e}" res = False - if params.do_diarization and not config.CAN_DIAR: + if params.do_diarization and not settings.CAN_DIAR: error_description += "Diarization is not available.\n" logger.error("Запрошена диаризация, но она не доступна.") params.do_diarization = False diff --git a/Recognizer/engine/sentensizer.py b/Recognizer/engine/sentensizer.py index ae9683c..42938a9 100644 --- a/Recognizer/engine/sentensizer.py +++ b/Recognizer/engine/sentensizer.py @@ -1,7 +1,7 @@ from utils.do_logging import logger import numpy as np import asyncio -import config +from config import settings from Punctuation import sbertpunc async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): @@ -81,7 +81,7 @@ async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): between_words_delta.append(word.get('end') - end_time) end_time = word.get('end') - words_mean = np.percentile(between_words_delta, config.BETWEEN_WORDS_PERCENTILE) + words_mean = np.percentile(between_words_delta, settings.BETWEEN_WORDS_PERCENTILE) logger.debug(f"words_mean = {words_mean}") start_time = 0 diff --git a/Recognizer/engine/stream_recognition.py b/Recognizer/engine/stream_recognition.py index 13d4dd5..b4ddc2d 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -1,6 +1,6 @@ import time import numpy as np -import config +from config import settings from Recognizer import recognizer from utils.bytes_to_samples_audio import get_np_array_samples_float32 from utils.resamppling import sync_resample_audiosegment @@ -43,12 +43,12 @@ def calc_speed(data): async def simple_recognise(audio_data, ) -> dict: # Приводим фреймрейт к фреймрейту модели - if audio_data.frame_rate != config.BASE_SAMPLE_RATE: - audio_data = sync_resample_audiosegment(audio_data, config.BASE_SAMPLE_RATE) + if audio_data.frame_rate != settings.BASE_SAMPLE_RATE: + audio_data = sync_resample_audiosegment(audio_data, settings.BASE_SAMPLE_RATE) # Перевод в семплы для распознавания. samples = get_np_array_samples_float32(audio_data.raw_data, audio_data.sample_width) - result = asdict(recognizer.recognize(samples, sample_rate=config.BASE_SAMPLE_RATE)) + result = asdict(recognizer.recognize(samples, sample_rate=settings.BASE_SAMPLE_RATE)) return result @@ -76,12 +76,12 @@ async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_sl if can_slow_down and multiplier == 1: speed = calc_speed(result) logger.debug(f"Скорость аудио {speed} единиц в секунду") - if speed > config.SPEECH_PER_SEC_NORM_RATE: - # print(max((config.SPEECH_PER_SEC_NORM_RATE - 1) / speed, 0.8)) + if speed > settings.SPEECH_PER_SEC_NORM_RATE: + # print(max((settings.SPEECH_PER_SEC_NORM_RATE - 1) / speed, 0.8)) result, speed, multiplier = await recognise_w_speed_correction(audio_data=audio_data, can_slow_down=True, - multiplier=max((config.SPEECH_PER_SEC_NORM_RATE-1)/speed, 0.8) + multiplier=max((settings.SPEECH_PER_SEC_NORM_RATE-1)/speed, 0.8) ) return result, speed, multiplier @@ -97,7 +97,7 @@ def simple_recognise_batch(list_audio_data: list, batch_size: int = 8) -> list: # и выравнивание по длине list_of_padded_samples = [samples_padding(samples_to_pad) for samples_to_pad in list_of_all_samples] - waveform = *pad_list(list_of_padded_samples), config.BASE_SAMPLE_RATE + waveform = *pad_list(list_of_padded_samples), settings.BASE_SAMPLE_RATE resampled = recognizer.resampler(*waveform) # Делаем препроцессинг на все данные сразу diff --git a/VoiceActivityDetector/__init__.py b/VoiceActivityDetector/__init__.py index 751633b..7dc52c8 100644 --- a/VoiceActivityDetector/__init__.py +++ b/VoiceActivityDetector/__init__.py @@ -1,4 +1,4 @@ -import config +from config import settings from utils.pre_start_init import paths from utils.do_logging import logger import requests @@ -30,7 +30,7 @@ raise FileExistsError else: from .do_vad import SileroVAD - vad=SileroVAD(paths.get("vad_model_path"), use_gpu= config.VAD_WITH_GPU) + vad=SileroVAD(paths.get("vad_model_path"), use_gpu= settings.VAD_WITH_GPU) # Нужно наблюдать за результатом работы в многопотоке (если он будет) # Если будут сбои, то переводить создание класса в отдельный процесс. - vad.set_mode(config.VAD_SENSITIVITY) + vad.set_mode(settings.VAD_SENSITIVITY) diff --git a/config.py b/config.py index f5968ac..817d735 100644 --- a/config.py +++ b/config.py @@ -34,10 +34,10 @@ class Settings(BaseSettings): # Logger settings LOGGING_LEVEL: str = 'DEBUG' LOGGING_FORMAT: str = '#%(levelname)-8s %(filename)s [LINE:%(lineno)d] [%(asctime)s] %(message)s' - FILENAME: str = Field(default_factory=lambda: f'logs/ASR-{date.today()}.log') + FILENAME: str = Field(default_factory=lambda: f'logs/ASR-started_{date.today()}') FILEMODE: str = 'a' # Срок хранения логов в днях - LOG_BACKUP_COUNT: int = 180 + LOG_BACKUP_COUNT: int = 60 IS_PROD: bool = True # Recognition settings @@ -150,40 +150,37 @@ def _compute_derived(self): # Все остальные модули проекта могут продолжать использовать `config.HOST`, # `config.PORT` и т.д. без изменений. # ============================================================================== -HOST = settings.HOST -PORT = settings.PORT -MODEL_NAME = settings.MODEL_NAME -BASE_SAMPLE_RATE = settings.BASE_SAMPLE_RATE -PROVIDER = settings.PROVIDER -NUM_THREADS = settings.NUM_THREADS -HF_HOME = settings.HF_HOME -LOGGING_LEVEL = settings.LOGGING_LEVEL -LOGGING_FORMAT = settings.LOGGING_FORMAT -FILENAME = settings.FILENAME -FILEMODE = settings.FILEMODE -LOG_BACKUP_COUNT = settings.LOG_BACKUP_COUNT -IS_PROD = settings.IS_PROD -MAX_OVERLAP_DURATION = settings.MAX_OVERLAP_DURATION -RECOGNITION_ATTEMPTS = settings.RECOGNITION_ATTEMPTS -SPEECH_PER_SEC_NORM_RATE = settings.SPEECH_PER_SEC_NORM_RATE -MAKE_MONO = settings.MAKE_MONO -USE_BATCH = settings.USE_BATCH -ASR_BATCH_SIZE = settings.ASR_BATCH_SIZE -VAD_SENSITIVITY = settings.VAD_SENSITIVITY -VAD_WITH_GPU = settings.VAD_WITH_GPU -BETWEEN_WORDS_PERCENTILE = settings.BETWEEN_WORDS_PERCENTILE -CAN_PUNCTUATE = settings.CAN_PUNCTUATE -PUNCTUATE_WITH_GPU = settings.PUNCTUATE_WITH_GPU -CAN_DIAR = settings.CAN_DIAR -DIAR_MODEL_NAME = settings.DIAR_MODEL_NAME -DIAR_WITH_GPU = settings.DIAR_WITH_GPU -CPU_WORKERS = settings.CPU_WORKERS -DIAR_GPU_BATCH_SIZE = settings.DIAR_GPU_BATCH_SIZE -DO_SPEED_SPEECH_CORRECTION = settings.DO_SPEED_SPEECH_CORRECTION -SPEED_SPEECH_CORRECTION_MULTIPLIER = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER -DO_LOCAL_FILE_RECOGNITIONS = settings.DO_LOCAL_FILE_RECOGNITIONS -DELETE_LOCAL_FILE_AFTR_ASR = settings.DELETE_LOCAL_FILE_AFTR_ASR -HUMAN_FORMAT_MD_FILE = settings.HUMAN_FORMAT_MD_FILE +# BASE_SAMPLE_RATE = settings.BASE_SAMPLE_RATE +# PROVIDER = settings.PROVIDER +# NUM_THREADS = settings.NUM_THREADS +# HF_HOME = settings.HF_HOME +# LOGGING_LEVEL = settings.LOGGING_LEVEL +# LOGGING_FORMAT = settings.LOGGING_FORMAT +# FILENAME = settings.FILENAME +# FILEMODE = settings.FILEMODE +# LOG_BACKUP_COUNT = settings.LOG_BACKUP_COUNT +# IS_PROD = settings.IS_PROD +# MAX_OVERLAP_DURATION = settings.MAX_OVERLAP_DURATION +# RECOGNITION_ATTEMPTS = settings.RECOGNITION_ATTEMPTS +# SPEECH_PER_SEC_NORM_RATE = settings.SPEECH_PER_SEC_NORM_RATE +# MAKE_MONO = settings.MAKE_MONO +# USE_BATCH = settings.USE_BATCH +# ASR_BATCH_SIZE = settings.ASR_BATCH_SIZE +# VAD_SENSITIVITY = settings.VAD_SENSITIVITY +# VAD_WITH_GPU = settings.VAD_WITH_GPU +# BETWEEN_WORDS_PERCENTILE = settings.BETWEEN_WORDS_PERCENTILE +# CAN_PUNCTUATE = settings.CAN_PUNCTUATE +# PUNCTUATE_WITH_GPU = settings.PUNCTUATE_WITH_GPU +# CAN_DIAR = settings.CAN_DIAR +# DIAR_MODEL_NAME = settings.DIAR_MODEL_NAME +# DIAR_WITH_GPU = settings.DIAR_WITH_GPU +# CPU_WORKERS = settings.CPU_WORKERS +# DIAR_GPU_BATCH_SIZE = settings.DIAR_GPU_BATCH_SIZE +# DO_SPEED_SPEECH_CORRECTION = settings.DO_SPEED_SPEECH_CORRECTION +# SPEED_SPEECH_CORRECTION_MULTIPLIER = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER +# DO_LOCAL_FILE_RECOGNITIONS = settings.DO_LOCAL_FILE_RECOGNITIONS +# DELETE_LOCAL_FILE_AFTR_ASR = settings.DELETE_LOCAL_FILE_AFTR_ASR +# HUMAN_FORMAT_MD_FILE = settings.HUMAN_FORMAT_MD_FILE AUDIOEXTENTIONS = [ # Основные форматы diff --git a/main.py b/main.py index f9b5702..5a6f5ae 100644 --- a/main.py +++ b/main.py @@ -1,6 +1,6 @@ from utils.do_logging import logger import uvicorn -import config +from config import settings import os import gc from contextlib import asynccontextmanager @@ -30,7 +30,7 @@ async def lifespan(app): # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) - if config.DO_LOCAL_FILE_RECOGNITIONS: + if settings.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), daemon=True @@ -108,7 +108,7 @@ def custom_openapi(): try: if __name__ == '__main__': app.openapi = custom_openapi - uvicorn.run(app, host=config.HOST, port=config.PORT) + uvicorn.run(app, host=settings.HOST, port=settings.PORT) except KeyboardInterrupt: logger.info('\nDone') except Exception as e: diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 2c34455..92212a7 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -2,7 +2,7 @@ from typing import Union, Annotated, Optional, Any, List, Dict from fastapi import UploadFile -import config +from config import settings class BaseResponse(BaseModel): @@ -65,10 +65,10 @@ class SyncASRRequest(BaseModel): do_diarization: Union[bool, None] = False make_mono: Union[bool, None] = False diar_vad_sensity: int = 3 - do_auto_speech_speed_correction: Union[bool, None] = config.DO_SPEED_SPEECH_CORRECTION - speech_speed_correction_multiplier: Union[float, None] = config.SPEED_SPEECH_CORRECTION_MULTIPLIER - use_batch: Union[bool, None] = config.USE_BATCH - batch_size: Union[int, None] = config.ASR_BATCH_SIZE + do_auto_speech_speed_correction: Union[bool, None] = settings.DO_SPEED_SPEECH_CORRECTION + speech_speed_correction_multiplier: Union[float, None] = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER + use_batch: Union[bool, None] = settings.USE_BATCH + batch_size: Union[int, None] = settings.ASR_BATCH_SIZE model_config = ConfigDict( json_schema_extra={ @@ -105,11 +105,11 @@ class PostFileRequest(BaseModel): do_dialogue: Union[bool, None] = False do_punctuation: Union[bool, None] = False do_diarization: Union[bool, None] = False - use_batch: Union[bool, None] = config.USE_BATCH - batch_size: Union[int, None] = config.ASR_BATCH_SIZE + use_batch: Union[bool, None] = settings.USE_BATCH + batch_size: Union[int, None] = settings.ASR_BATCH_SIZE diar_vad_sensity: int = 3 - do_auto_speech_speed_correction: Union[bool, None] = config.DO_SPEED_SPEECH_CORRECTION - speech_speed_correction_multiplier: Union[float, None] = config.SPEED_SPEECH_CORRECTION_MULTIPLIER + do_auto_speech_speed_correction: Union[bool, None] = settings.DO_SPEED_SPEECH_CORRECTION + speech_speed_correction_multiplier: Union[float, None] = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER make_mono: Union[bool, None] = False model_config = ConfigDict( diff --git a/routes/is_alive.py b/routes/is_alive.py index d1735e2..07f745a 100644 --- a/routes/is_alive.py +++ b/routes/is_alive.py @@ -4,6 +4,7 @@ import os import pynvml from utils.pre_start_init import audio_to_asr +from models.fast_api_models import BaseResponse router = APIRouter() @@ -24,7 +25,7 @@ def get_gpu_free_memory(): return free_mb, gpu_load,temperature -@router.get("/is_alive") +@router.get("/is_alive", response_model=BaseResponse) async def check_if_service_is_alive(): logging.info('GET_is_alive') @@ -37,11 +38,14 @@ async def check_if_service_is_alive(): else: state = "in_work" - return {"error": False, - "error_description": None, + return BaseResponse( + success=True, + error_description=None, + data={ "state": state, "tasks_in_work": tasks_in_work, "free_memory_mb": free_mb, "gpu_load_percent": gpu_load, "temperature_celsius": temperature - } + } + ) diff --git a/routes/post_by_file_FORM.py b/routes/post_by_file_FORM.py index 4e9aba8..413a92e 100644 --- a/routes/post_by_file_FORM.py +++ b/routes/post_by_file_FORM.py @@ -1,7 +1,7 @@ from io import BytesIO import asyncio -import config +from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile from utils.do_logging import logger from models.fast_api_models import PostFileRequest @@ -23,8 +23,8 @@ def get_file_request( do_punctuation: bool = Form(default=False, description="Восстанавливать пунктуацию."), do_diarization: bool = Form(default=False, description="Разделять по спикерам."), diar_vad_sensity: int = Form(default=3, description="Чувствительность VAD."), - use_batch: bool = Form(default=config.USE_BATCH, description="Использовать батчинг для ASR."), - batch_size: int = Form(default=config.ASR_BATCH_SIZE, description="Размер батча для ASR."), + use_batch: bool = Form(default=settings.USE_BATCH, description="Использовать батчинг для ASR."), + batch_size: int = Form(default=settings.ASR_BATCH_SIZE, description="Размер батча для ASR."), do_auto_speech_speed_correction: bool = Form(default=True, description="Корректировать скорость речи при распознавании."), speech_speed_correction_multiplier: float = Form(default=1, description="Базовый коэффициент скорости речи."), make_mono: bool = Form(default=False, description="Соединить несколько каналов в mono"), diff --git a/routes/post_by_url.py b/routes/post_by_url.py index 5305658..18f1aac 100644 --- a/routes/post_by_url.py +++ b/routes/post_by_url.py @@ -5,7 +5,7 @@ from utils.pre_start_init import posted_and_downloaded_audio from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile -from models.fast_api_models import SyncASRRequest +from models.fast_api_models import SyncASRRequest, BaseResponse from Recognizer.engine.file_recognition import process_file from threading import Lock from io import BytesIO diff --git a/routes/root.py b/routes/root.py index 73f7cfd..32ec78f 100644 --- a/routes/root.py +++ b/routes/root.py @@ -1,13 +1,22 @@ from fastapi import APIRouter +from models.fast_api_models import BaseResponse import config router = APIRouter() -@router.get("/") +@router.get("/", response_model=BaseResponse) async def root(): - print("Зашли в root") - - return {"error": True, - "data": "No_service_selected", - # "available_services": ["Vosk_Recognizer"], - "comment": f"try_addr: http://{config.HOST}:{config.PORT}/docs"} + return BaseResponse( + success=True, + error_description=None, + data={"message": "No_service_selected", + "available_endpoints": { + "POST /post_one_step_req": "ASR by URL", + "POST /post_file": "ASR by file upload", + "WS /ws": "WebSocket streaming ASR", + "GET /is_alive": "Service health check", + "GET /docs": "API documentation", + "/demo": "DEMO UI page" + }, + "try_addr": f"http://{config.settings.HOST}:{config.settings.PORT}/docs"} + ) diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 8b9b221..0458a32 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -1,7 +1,7 @@ from pydub import AudioSegment import ujson -import config +from config import settings import uuid from io import BytesIO @@ -24,15 +24,15 @@ async def websocket(ws: WebSocket): wait_null_answers=True client_id = uuid.uuid4() logger.debug(f'Принят новый сокет id = {client_id}') - audio_buffer[client_id] = AudioSegment.silent(1, frame_rate=config.BASE_SAMPLE_RATE) - audio_overlap[client_id] = AudioSegment.silent(1, frame_rate=config.BASE_SAMPLE_RATE) + audio_buffer[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + audio_overlap[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) audio_duration[client_id] = 0 audio_to_asr[client_id] = list() ws_collected_asr_res[client_id] = {f"channel_{1}": list()} do_dialogue = False do_punctuation = False audio_format = 'raw' - sample_rate = config.BASE_SAMPLE_RATE # Если не получен фреймрейт в конфиге сокета, по попытается принять с конфигом модели. + sample_rate = settings.BASE_SAMPLE_RATE # Если не получен фреймрейт в конфиге сокета, по попытается принять с конфигом модели. sentenced_data = None error_description = None @@ -97,8 +97,8 @@ async def websocket(ws: WebSocket): logger.debug(f"Чанк принят и распознан in channel {channel_name}") # Приводим фреймрейт к фреймрейту модели - if audiosegment_chunk.frame_rate != config.BASE_SAMPLE_RATE: - audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, config.BASE_SAMPLE_RATE) + if audiosegment_chunk.frame_rate != settings.BASE_SAMPLE_RATE: + audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, settings.BASE_SAMPLE_RATE) if audiosegment_chunk.channels != 1: audiosegment_chunk = audiosegment_chunk.set_channels(1) @@ -107,7 +107,7 @@ async def websocket(ws: WebSocket): audio_buffer[client_id] += audiosegment_chunk # Накопили больше нормы - if (audio_overlap[client_id]+audio_buffer[client_id]).duration_seconds >= config.MAX_OVERLAP_DURATION: + if (audio_overlap[client_id]+audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: # Проверяем новый чанк перед объединением (там же режем хвост и добавляем его при необходимости) await find_last_speech_position(client_id, is_last_chunk=False) else: diff --git a/utils/bytes_to_samples_audio.py b/utils/bytes_to_samples_audio.py index b82faf5..a52e685 100644 --- a/utils/bytes_to_samples_audio.py +++ b/utils/bytes_to_samples_audio.py @@ -1,8 +1,7 @@ import numpy as np from time import perf_counter import logging - -import config +from config import settings def get_np_array_samples_float32(audio_bytes: bytes, sample_width: int = 2) -> np.ndarray: @@ -58,7 +57,7 @@ def get_np_array_samples_float32(audio_bytes: bytes, sample_width: int = 2) -> n f"Длина samples_float32: {len(samples_float32)}, min={np.min(samples_float32)}, max={np.max(samples_float32)}") end_timer = perf_counter() - logger.debug(f"Время на конвертацию аудио длинной {len(samples)/config.BASE_SAMPLE_RATE} сек. затрачено {(end_timer-timer_start)} сек.") + logger.debug(f"Время на конвертацию аудио длинной {len(samples)/settings.BASE_SAMPLE_RATE} сек. затрачено {(end_timer-timer_start)} сек.") return samples_float32 except Exception as e: diff --git a/utils/chunk_doing.py b/utils/chunk_doing.py index 8402448..3d81f92 100644 --- a/utils/chunk_doing.py +++ b/utils/chunk_doing.py @@ -1,4 +1,4 @@ -import config +from config import settings import numpy as np from pydub import AudioSegment from utils.do_logging import logger @@ -25,11 +25,11 @@ async def find_last_speech_position(socket_id, is_last_chunk): if is_last_chunk: last_audio = audio_overlap[socket_id] + audio_buffer[socket_id] - # for i in last_audio[::config.MAX_OVERLAP_DURATION*1000]: + # for i in last_audio[::settings.MAX_OVERLAP_DURATION*1000]: # audio_to_asr[socket_id].append(last_audio[:i]) - for i in range(0, len(last_audio), config.MAX_OVERLAP_DURATION*1000): - audio_to_asr[socket_id].append(last_audio[i:min(i + config.MAX_OVERLAP_DURATION*1000, len(last_audio))]) + for i in range(0, len(last_audio), settings.MAX_OVERLAP_DURATION*1000): + audio_to_asr[socket_id].append(last_audio[i:min(i + settings.MAX_OVERLAP_DURATION*1000, len(last_audio))]) @@ -83,7 +83,7 @@ async def find_last_speech_position(socket_id, is_last_chunk): min_silence_frames = int(duration_seconds / frame_duration) # Устанавливаем стартовые значения. - max_audio_length = len(audio) if len(audio) < config.MAX_OVERLAP_DURATION*silero_bitrate else config.MAX_OVERLAP_DURATION*silero_bitrate + max_audio_length = len(audio) if len(audio) < settings.MAX_OVERLAP_DURATION*silero_bitrate else settings.MAX_OVERLAP_DURATION*silero_bitrate partial_frame_length = 0 @@ -148,7 +148,7 @@ async def find_last_speech_position(socket_id, is_last_chunk): return -def samples_padding(samples, sample_rate = config.BASE_SAMPLE_RATE, duration = config.MAX_OVERLAP_DURATION) -> np.ndarray: +def samples_padding(samples, sample_rate = settings.BASE_SAMPLE_RATE, duration = settings.MAX_OVERLAP_DURATION) -> np.ndarray: max_samples_len = int(duration * sample_rate) # Выравнивание до максимальной длины @@ -158,8 +158,8 @@ def samples_padding(samples, sample_rate = config.BASE_SAMPLE_RATE, duration = c padded_samples[:len(samples)] = samples elif len(samples) > max_samples_len: # Обрезка до максимальной длины - logger.warning(f"Аудио длиной {len(samples) / config.BASE_SAMPLE_RATE:.2f} сек. " - f"превышает MAX_OVERLAP_DURATION ({config.MAX_OVERLAP_DURATION} сек.). " + logger.warning(f"Аудио длиной {len(samples) / settings.BASE_SAMPLE_RATE:.2f} сек. " + f"превышает MAX_OVERLAP_DURATION ({settings.MAX_OVERLAP_DURATION} сек.). " f"Будет обрезано до {max_samples_len} семплов.") padded_samples = samples[:max_samples_len] diff --git a/utils/do_logging.py b/utils/do_logging.py index 8b4b271..ff418a5 100644 --- a/utils/do_logging.py +++ b/utils/do_logging.py @@ -1,34 +1,35 @@ # -*- coding: utf-8 -*- -import config +from config import settings + import logging from logging.handlers import TimedRotatingFileHandler from fastapi.logger import logger as fastapi_logger logger = logging.getLogger(__name__) -if config.IS_PROD: +if settings.IS_PROD: # Создаем обработчик, который ротирует логи каждый день в полночь file_handler = TimedRotatingFileHandler( - filename=config.FILENAME, # Базовое имя файла + filename=settings.FILENAME, # Базовое имя файла when='midnight', # Ротация каждый день в полночь interval=1, # Интервал - каждый день - backupCount=config.LOG_BACKUP_COUNT if hasattr(config, 'LOG_BACKUP_COUNT') else 7, + backupCount=settings.LOG_BACKUP_COUNT if hasattr(settings, 'LOG_BACKUP_COUNT') else 7, # Хранить 7 дней логов по умолчанию encoding='UTF-8' ) - file_handler.setLevel(config.LOGGING_LEVEL) - file_handler.setFormatter(logging.Formatter(config.LOGGING_FORMAT)) + file_handler.setLevel(settings.LOGGING_LEVEL) + file_handler.setFormatter(logging.Formatter(settings.LOGGING_FORMAT)) # Настраиваем логгер logger.addHandler(file_handler) fastapi_logger.addHandler(file_handler) # Убираем basicConfig, так как мы используем кастомный обработчик - logging.basicConfig(level=config.LOGGING_LEVEL) + logging.basicConfig(level=settings.LOGGING_LEVEL) else: logging.basicConfig( - level=config.LOGGING_LEVEL, - format=config.LOGGING_FORMAT, + level=settings.LOGGING_LEVEL, + format=settings.LOGGING_FORMAT, encoding="UTF-8" ) \ No newline at end of file diff --git a/utils/files_whatcher.py b/utils/files_whatcher.py index a871046..7b93208 100644 --- a/utils/files_whatcher.py +++ b/utils/files_whatcher.py @@ -6,7 +6,8 @@ from watchdog.events import FileSystemEventHandler, DirMovedEvent, FileMovedEvent from pathlib import Path -import config +from config import settings + from models.fast_api_models import PostFileRequest from utils.save_asr_to_md import save_to_file @@ -16,11 +17,11 @@ def send_file_to_asr(event, file_path): if not event.is_directory: file_params = PostFileRequest() file_params.speech_speed_correction_multiplier = 1 - file_params.do_diarization = config.CAN_DIAR - file_params.do_punctuation = config.CAN_PUNCTUATE + file_params.do_diarization = settings.CAN_DIAR + file_params.do_punctuation = settings.CAN_PUNCTUATE file_params.do_dialogue = True file_params.do_echo_clearing = True - file_params.make_mono = config.MAKE_MONO + file_params.make_mono = settings.MAKE_MONO file_params.diar_vad_sensity = 2 file_params.do_auto_speech_speed_correction = True file_params.keep_raw = False @@ -35,7 +36,7 @@ def send_file_to_asr(event, file_path): if asr_data["success"]: asyncio.run(save_to_file(asr_data, file_path)) - if config.DELETE_LOCAL_FILE_AFTR_ASR: + if settings.DELETE_LOCAL_FILE_AFTR_ASR: try: file_path.unlink() logging.info(f"Файл {file_path} удалён") @@ -54,7 +55,7 @@ def on_created(self, event): # Todo - перенести paths в отдельный файл и импортировать его по необходимости file_path = Path(event.src_path) logging.info(f"Получено сообщение о новом файле {event.src_path}") - if file_path.suffix not in config.AUDIOEXTENTIONS: + if file_path.suffix not in settings.AUDIOEXTENTIONS: logging.info(f"Файл {file_path.name} пропущен, т.к. не аудио формат") else: send_file_to_asr(event,file_path) @@ -63,7 +64,7 @@ def on_moved(self, event: DirMovedEvent | FileMovedEvent) -> None: file_path = Path(event.dest_path) logging.info(f"Получено сообщение о переименовании файла {event.dest_path}") - if file_path.suffix not in config.AUDIOEXTENTIONS: + if file_path.suffix not in settings.AUDIOEXTENTIONS: logging.info(f"Файл {file_path.name} пропущен, т.к. не аудио формат") else: send_file_to_asr(event, file_path) diff --git a/utils/pre_start_init.py b/utils/pre_start_init.py index e47a069..11a7c54 100644 --- a/utils/pre_start_init.py +++ b/utils/pre_start_init.py @@ -1,7 +1,6 @@ # -*- coding: utf-8 -*- from pathlib import Path -import config - +from config import settings BASE_DIR = Path(__file__).resolve().parent.parent paths = { @@ -34,7 +33,7 @@ "punctuation_model_path": BASE_DIR / "models" / "sbert_punc_case_ru_onnx", "vad_model_path": BASE_DIR / "models" / "VAD_silero_v5" / "silero_vad.onnx", - "diar_speaker_model_path": BASE_DIR / "models" / "DIARISATION_model" / f"{config.DIAR_MODEL_NAME}", + "diar_speaker_model_path": BASE_DIR / "models" / "DIARISATION_model" / f"{settings.DIAR_MODEL_NAME}", "BASE_DIR": BASE_DIR, "test_file": BASE_DIR /'trash'/'111.wav', From 25d5d9ff4bc039905d0694c4ac29c297df71e075 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Tue, 28 Apr 2026 17:14:14 +0300 Subject: [PATCH 10/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20v1=20API=20=D0=B4=D0=BB=D1=8F=20?= =?UTF-8?q?=D1=81=D0=BE=D1=85=D1=80=D0=B0=D0=BD=D0=B5=D0=BD=D0=B8=D1=8F=20?= =?UTF-8?q?=D1=84=D1=83=D0=BD=D0=BA=D1=86=D0=B8=D0=BE=D0=BD=D0=B0=D0=BB?= =?UTF-8?q?=D0=B0=20=D0=B8=20=D0=BF=D1=80=D0=B8=D0=B2=D0=B5=D0=B4=D0=B5?= =?UTF-8?q?=D0=BD=D0=B8=D1=8F=20=D0=BE=D1=82=D0=B2=D0=B5=D1=82=D0=BE=D0=B2?= =?UTF-8?q?=20=D0=BA=20=D0=B5=D0=B4=D0=B8=D0=BD=D0=BE=D0=B9=20=D1=81=D1=82?= =?UTF-8?q?=D1=80=D1=83=D0=BA=D1=82=D1=83=D1=80=D0=B5.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- global_to_do.txt | 151 --------------------------------- main.py | 55 +++++++++++- models/fast_api_models.py | 18 +++- routes/is_alive.py | 22 ++--- routes/post_by_file_FORM.py | 38 ++++----- routes/post_by_url.py | 54 ++++++------ routes/v1/__init__.py | 9 ++ routes/v1/is_alive.py | 33 +++++++ routes/v1/post_by_file_FORM.py | 51 +++++++++++ routes/v1/post_by_url.py | 46 ++++++++++ tests/__init__.py | 0 tests/test_http_routes.py | 105 +++++++++++++++++++++++ 12 files changed, 373 insertions(+), 209 deletions(-) delete mode 100644 global_to_do.txt create mode 100644 routes/v1/__init__.py create mode 100644 routes/v1/is_alive.py create mode 100644 routes/v1/post_by_file_FORM.py create mode 100644 routes/v1/post_by_url.py create mode 100644 tests/__init__.py create mode 100644 tests/test_http_routes.py diff --git a/global_to_do.txt b/global_to_do.txt deleted file mode 100644 index 61fdeeb..0000000 --- a/global_to_do.txt +++ /dev/null @@ -1,151 +0,0 @@ -================================================================================ - ПЛАН РЕФАКТОРИНГА ASR FastAPI - "Одна задача за раз" -================================================================================ - -ПРИНЦИПЫ РАБОТЫ: - - Одна задача за раз. - - Жди подтверждения перед правками. - - Сначала описывается план, пользователь соглашается — потом выполняется. - - Не делать несколько изменений сразу. - - После каждого этапа требуется тестирование перед переходом к следующему. - -================================================================================ -ЭТАП 1. ФУНДАМЕНТ: КОНФИГУРАЦИЯ И ТОЧКА ВХОДА -================================================================================ - -Цель: устранить side-effects при импорте, валидировать окружение и централизовать -создание приложения. - -ЗАДАЧА 1.1 — config.py [ВЫПОЛНЕНО] - - Миграция с голого os.getenv() на pydantic-settings (BaseSettings). - - Все переменные получают типизацию, значения по умолчанию и валидацию. - - Старые имена экспортируются для обратной совместимости с неизменяемыми - модулями ASR (Recognizer, Diarisation и т.д.). - - Добавить SettingsConfigDict с поддержкой .env файла. - -ЗАДАЧА 1.2 — utils/pre_start_init.py [ВЫПОЛНЕНО] - - Удалить глобальный объект app = FastAPI(...). - - Оставить только инициализацию paths, глобальных defaultdict - (audio_overlap, audio_buffer, audio_to_asr, audio_duration, - ws_collected_asr_res, posted_and_downloaded_audio) и lifespan. - - Убрать импорт FastAPI из этого файла. - - Убрать любие импорты роутов или HTTP-объектов. - - Файл должен быть безопасен для импорта без побочных эффектов. - -ЗАДАЧА 1.3 — main.py [ВЫПОЛНЕНО] - - Создание app = FastAPI(...) переносится сюда как явная фабрика. - - Подключение lifespan из utils.pre_start_init. - - Добавление CORS-middleware (CORSMiddleware). - - Перенос монтирования статики (/static) из routes/demo_page.py сюда. - - Точка входа if __name__ == '__main__': остается здесь. - -================================================================================ -ЭТАП 2. МАРШРУТИЗАЦИЯ: ОТКАЗ ОТ ГЛОБАЛЬНОГО app -================================================================================ - -Цель: разорвать жесткую связность роутов, убрать циклические импорты -и подготовить версионирование API. - -ЗАДАЧА 2.1 — routes/root.py, routes/is_alive.py, routes/post_ws.py - - Заменить импорт глобального app на создание APIRouter() в каждом файле. - - Декораторы меняются с @app.* на @router.*. - - Для is_alive добавить prefix="/is_alive" и tags=["healthcheck"]. - - Для post_ws добавить prefix="/ws". - -ЗАДАЧА 2.2 — routes/post_by_file_FORM.py, routes/post_by_url.py - - Аналогичный перевод на APIRouter. - - Добавить tags=["ASR"]. - - Убрать импорт app из utils.pre_start_init. - -ЗАДАЧА 2.3 — routes/ws_audio_transkrib.py, routes/demo_page.py - - Перевод WebSocket и HTML-роута на APIRouter. - - Очистка неиспользуемых импортов (subprocess, WebSocketException). - - Убрать app.mount("/static", ...) из demo_page.py (перенесено в main.py). - -После завершения этого этапа в main.py появляется единственное место -подключения всех роутеров через app.include_router(...). - -================================================================================ -ЭТАП 3. МОДЕЛИ API: ЕДИНЫЙ КОНТРАКТ ОТВЕТОВ -================================================================================ - -Цель: стандартизировать формат ответов API и заложить структуру для будущей -авторизации. - -ЗАДАЧА 3.1 — models/fast_api_models.py - - Добавить базовые Pydantic-модели: - * BaseResponse / ErrorResponse (унификация полей success, - error_description, data). - * UserBase, Token, TokenPayload — заготовки для JWT-аутентификации. - - Реорганизовать существующие SyncASRRequest, PostFileRequest, - PostFileRequestDiarize, WebSocketModel без потери функциональности. - - Добавить примеры ответов (json_schema_extra) для документации. - -================================================================================ -ЭТАП 4. ИНФРАСТРУКТУРА HIGHLOAD (ПОДГОТОВКА) -================================================================================ - -Цель: подготовить приложение к работе за reverse-proxy, в k8s и под нагрузкой. - -ЗАДАЧА 4.1 — main.py (дополнение) - - Добавить middleware: - * TrustedHostMiddleware. - * Обработчик ошибок (HTTPException -> JSON-ответ). - * Генерация request_id для трейсинга (или через middleware, - или через correlation_id). - - Настроить корректные заголовки для проксирования (X-Forwarded-For). - -ЗАДАЧА 4.2 — routes/is_alive.py (расширение) - - Разделить эндпоинт на: - * /health/live (liveness probe) — проверка, что процесс жив. - * /health/ready (readiness probe) — проверка готовности принимать - трафик (модели загружены, память GPU в норме). - - Добавить метрики: uptime, количество обработанных запросов (заготовка). - -ЗАДАЧА 4.3 — создание core/logging_config.py (новый файл) - - Заготовка структурированного JSON-логирования. - - Формат: {"timestamp": "...", "level": "...", "message": "...", - "request_id": "...", "module": "..."}. - - Подготовка для последующей интеграции с ELK / Loki / Grafana. - -================================================================================ -ЭТАП 5. ПОДГОТОВКА К AUTH, ADMIN, PAYMENTS -================================================================================ - -Цель: создать слой security и структуру под будущие модули, не ломая текущий -ASR-функционал. - -ЗАДАЧА 5.1 — создание core/security.py (новый файл) - - Утилиты для хеширования паролей (passlib / bcrypt). - - Создание и верификация JWT-токенов (jose / PyJWT). - - SECRET_KEY, ALGORITHM, ACCESS_TOKEN_EXPIRE_MINUTES из config. - -ЗАДАЧА 5.2 — создание api/deps.py (новый файл) - - Зависимости FastAPI (Depends): - * get_current_user — извлечение пользователя из токена. - * get_current_active_user — проверка статуса пользователя. - - Пока реализовать как заглушки (stub), возвращающие фиктивного пользователя. - -ЗАДАЧА 5.3 — создание структуры api/v1/ (новые директории) - - Перенос текущих роутов под префикс /api/v1/. - - Освобождение корневых путей (/, /demo, /docs) для админки и фронта. - - Создание api/v1/endpoints/ для ASR-роутов. - - Создание api/v1/api.py для агрегации роутеров v1. - -ЗАДАЧА 5.4 — models/fast_api_models.py (дополнение) - - Добавить модели для платежной системы (заготовки): - * Subscription (тип подписки, дата начала/окончания, статус). - * Transaction (id, сумма, валюта, статус, внешний id платежа). - * UserAccount (баланс, тариф). - - Добавить enum-ы для статусов. - -================================================================================ - ПОРЯДОК СТАРТА -================================================================================ - -Стартовать рекомендуется с Этапа 1, Задача 1.1 (config.py). -После согласования каждой задачи выдается полный текст измененного файла -(или согласованной группы файлов) для замены в репозитории. - -================================================================================ diff --git a/main.py b/main.py index 5a6f5ae..73c0bcc 100644 --- a/main.py +++ b/main.py @@ -4,7 +4,10 @@ import os import gc from contextlib import asynccontextmanager -from fastapi import FastAPI +from fastapi import FastAPI, Request +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse +from starlette.exceptions import HTTPException as StarletteHTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles from fastapi.openapi.utils import get_openapi @@ -19,6 +22,8 @@ from routes.post_by_url import router as post_by_url_router from routes.ws_audio_transkrib import router as ws_audio_transkrib_router from routes.demo_page import router as demo_router +from routes.v1 import router as v1_router +from models.fast_api_models import ErrorResponse import models @@ -30,6 +35,9 @@ async def lifespan(app): # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) + # Установка HF_HOME для HuggingFace Hub + os.environ["HF_HOME"] = settings.HF_HOME + if settings.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), @@ -49,6 +57,50 @@ async def lifespan(app): title='ASR' ) + +@app.exception_handler(RequestValidationError) +async def validation_exception_handler(request: Request, exc: RequestValidationError): + return JSONResponse( + status_code=422, + content=ErrorResponse( + success=False, + error_description=str(exc), + raw_data={}, + sentenced_data={}, + diarized_data={}, + ).model_dump() + ) + + +@app.exception_handler(StarletteHTTPException) +async def http_exception_handler(request: Request, exc: StarletteHTTPException): + return JSONResponse( + status_code=exc.status_code, + content=ErrorResponse( + success=False, + error_description=exc.detail, + raw_data={}, + sentenced_data={}, + diarized_data={}, + ).model_dump() + ) + + +@app.exception_handler(Exception) +async def general_exception_handler(request: Request, exc: Exception): + logger.error(f"Unhandled exception: {exc}") + return JSONResponse( + status_code=500, + content=ErrorResponse( + success=False, + error_description="Internal server error", + raw_data={}, + sentenced_data={}, + diarized_data={}, + ).model_dump() + ) + + # CORS middleware app.add_middleware( CORSMiddleware, @@ -69,6 +121,7 @@ async def lifespan(app): app.include_router(post_by_url_router) app.include_router(ws_audio_transkrib_router) app.include_router(demo_router) +app.include_router(v1_router) def custom_openapi(): openapi_schema = get_openapi( diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 92212a7..f52cc0e 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -11,7 +11,19 @@ class BaseResponse(BaseModel): """ success: bool = True error_description: Optional[str] = None - data: Optional[Dict[str, Any]] = None + raw_data: Optional[Dict[str, Any]] = None + sentenced_data: Optional[Dict[str, Any]] = None + diarized_data: Optional[Dict[str, Any]] = None + + +class V1BaseResponse(BaseModel): + """ + Единый формат ответа для API v1. + Все данные помещаются в поле data. + """ + success: bool = True + error_description: Optional[str] = None + data: dict = {} class ErrorResponse(BaseModel): @@ -20,7 +32,9 @@ class ErrorResponse(BaseModel): """ success: bool = False error_description: str - data: Optional[Dict[str, Any]] = None + raw_data: Optional[Dict[str, Any]] = None + sentenced_data: Optional[Dict[str, Any]] = None + diarized_data: Optional[Dict[str, Any]] = None class UserBase(BaseModel): diff --git a/routes/is_alive.py b/routes/is_alive.py index 07f745a..afec665 100644 --- a/routes/is_alive.py +++ b/routes/is_alive.py @@ -11,6 +11,7 @@ def get_gpu_free_memory(): try: + # todo Доработать, выдавать ответ в зависимости от провайдера. pynvml.nvmlInit() handle = pynvml.nvmlDeviceGetHandleByIndex(0) # Первая видеокарта mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle) @@ -19,33 +20,32 @@ def get_gpu_free_memory(): gpu_load = utilization.gpu temperature = pynvml.nvmlDeviceGetTemperature(handle, pynvml.NVML_TEMPERATURE_GPU) except pynvml.NVMLError as e: - return {"error": str(e)} + return {"error": str(e)}, None, None, None finally: pynvml.nvmlShutdown() - return free_mb, gpu_load,temperature + return None, free_mb, gpu_load,temperature -@router.get("/is_alive", response_model=BaseResponse) +@router.get("/is_alive") async def check_if_service_is_alive(): - + error_description = None logging.info('GET_is_alive') tasks_in_work = len(audio_to_asr) - free_mb, gpu_load,temperature = get_gpu_free_memory() + error, free_mb, gpu_load,temperature = get_gpu_free_memory() + if error: + error_description = error.get("error",None) if tasks_in_work == 0: state = "idle" else: state = "in_work" - return BaseResponse( - success=True, - error_description=None, - data={ + return {"error": False, + "error_description": error_description, "state": state, "tasks_in_work": tasks_in_work, "free_memory_mb": free_mb, "gpu_load_percent": gpu_load, "temperature_celsius": temperature - } - ) + } diff --git a/routes/post_by_file_FORM.py b/routes/post_by_file_FORM.py index 413a92e..961f551 100644 --- a/routes/post_by_file_FORM.py +++ b/routes/post_by_file_FORM.py @@ -4,7 +4,7 @@ from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile from utils.do_logging import logger -from models.fast_api_models import PostFileRequest +from models.fast_api_models import PostFileRequest, BaseResponse from Recognizer.engine.file_recognition import process_file from threading import Lock @@ -44,21 +44,11 @@ def get_file_request( ) -@router.post("/post_file") +@router.post("/post_file", response_model=BaseResponse) async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), -): - res = True - error_description = str() - - result = { - "success": res, - "error_description": error_description, - "raw_data": dict(), - "sentenced_data": dict(), - } - +) -> BaseResponse: # Сохраняем файл на диск асинхронно try: buffer = BytesIO(await file.read()) @@ -66,19 +56,29 @@ async def async_receive_file( except Exception as e: error_description = f"Не удалось сохранить файл для распознавания: {file.filename}, размер файла: {file.size}, по причине: {e}" logger.error(error_description) - result["success"] = False - result["error_description"] = error_description - return result + return BaseResponse( + success=False, + error_description=error_description, + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) else: logger.info(f"Получен и сохранён файл {file.filename}") try: # Запускаем обработку в потоке - result = await asyncio.to_thread(process_file, buffer, params) + result_dict = await asyncio.to_thread(process_file, buffer, params) + result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) + result = BaseResponse( + success=False, + error_description=str(error_description), + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) finally: # Удаляем временный файл await file.close() diff --git a/routes/post_by_url.py b/routes/post_by_url.py index 18f1aac..fdc2987 100644 --- a/routes/post_by_url.py +++ b/routes/post_by_url.py @@ -16,27 +16,24 @@ # Глобальный лок для потокобезопасности audio_lock = Lock() -@router.post("/post_one_step_req") -async def post(params: SyncASRRequest): +@router.post("/post_one_step_req", response_model=BaseResponse) +async def post(params: SyncASRRequest) -> BaseResponse: """ - На вход принимает HttpUrl - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n + На вход ждёт str(HttpUrl) - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n Если на вход передаётся не моно, то ответ будет в несколько элементов списка для каждого канала.\n - По умолчанию отдаёт сырой результат распознавания с разбивкой на части продолжительностью около 15 секунд\n :param: do_dialogue: - true, если нужно разбить речь на диалог\n - :param: do_punctuation - true, если нужно расставить пунктуацию. Применяется к диалогу, общему тексту. В проекте.\n - При проектировании таймаутов учитывайте скорость распознавания (около 100 секунд аудио распознаётся за 2-5 секунд - распознавания одного канала) + :param: do_punctuation - true, если нужно расставить пунктуацию. Применяется к диалогу, и отдельно к общему тексту.\n + :param:keep_raw: Сохранять ли в выводе "сырые данные" - распознавание по словам. \n + :param:do_echo_clearing: Очищать от межканального эха \n + :param:do_diarization: Разделать речь на спикеров. Работает только с моно файлами. \n + :param:make_mono: Объединить каналы аудио в моно файл \n + :param:diar_vad_sensity: Чувствительность детектора голоса. \n + :param:do_auto_speech_speed_correction: Корректировать скорость речи (для очень быстрой речи). \n + :param:speech_speed_correction_multiplier: Задать коэффициент корректировки скорости речи \n + :param:use_batch: Union[bool, None] = Использовать пакетную обработку. Полезно при невозможности использовать Tensorrt \n + :param:batch_size: Union[int, None] = Размер пакета для обработки. \n """ - res = True - error_description = str() - - result = { - "success": res, - "error_description": error_description, - "raw_data": dict(), - "sentenced_data": dict(), - } # Получаем файл post_id = uuid.uuid4() @@ -47,20 +44,27 @@ async def post(params: SyncASRRequest): if not res: logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') - return { - "success": False, - "error_description": error_description, - "raw_data": dict(), - "sentenced_data": dict(), - } + return BaseResponse( + success=False, + error_description=error_description, + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) try: # Запускаем обработку в потоке - result = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) + result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) + result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) + return BaseResponse( + success=False, + error_description=str(error_description), + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) return result diff --git a/routes/v1/__init__.py b/routes/v1/__init__.py new file mode 100644 index 0000000..0fdb9b8 --- /dev/null +++ b/routes/v1/__init__.py @@ -0,0 +1,9 @@ +from fastapi import APIRouter +from .is_alive import router as is_alive_router +from .post_by_url import router as post_by_url_router +from .post_by_file_FORM import router as post_by_file_router + +router = APIRouter(prefix="/v1") +router.include_router(is_alive_router) +router.include_router(post_by_url_router) +router.include_router(post_by_file_router) diff --git a/routes/v1/is_alive.py b/routes/v1/is_alive.py new file mode 100644 index 0000000..a4db1d8 --- /dev/null +++ b/routes/v1/is_alive.py @@ -0,0 +1,33 @@ +from fastapi import APIRouter +import logging +from utils.pre_start_init import audio_to_asr +from routes.is_alive import get_gpu_free_memory +from models.fast_api_models import V1BaseResponse + +router = APIRouter() + +@router.get("/is_alive", response_model=V1BaseResponse) +async def check_if_service_is_alive_v1(): + logging.info('GET /v1/is_alive') + error_description = None + tasks_in_work = len(audio_to_asr) + error, free_mb, gpu_load, temperature = get_gpu_free_memory() + state = "idle" if tasks_in_work == 0 else "in_work" + + if error: + error_description = error.get("error",None) + data = None + else: + data = { + "state": state, + "tasks_in_work": tasks_in_work, + "free_memory_mb": free_mb, + "gpu_load_percent": gpu_load, + "temperature_celsius": temperature + } + + return V1BaseResponse( + success=True, + error_description=error_description, + data=data + ) diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py new file mode 100644 index 0000000..449558e --- /dev/null +++ b/routes/v1/post_by_file_FORM.py @@ -0,0 +1,51 @@ +from io import BytesIO +import asyncio +from fastapi import APIRouter, Depends, File, Form, UploadFile +from config import settings +from utils.do_logging import logger +from models.fast_api_models import PostFileRequest, V1BaseResponse +from Recognizer.engine.file_recognition import process_file +from routes.post_by_file_FORM import get_file_request + +router = APIRouter() + +@router.post("/post_file", response_model=V1BaseResponse) +async def async_receive_file_v1( + file: UploadFile = File(description="Аудиофайл для обработки"), + params: PostFileRequest = Depends(get_file_request), +) -> V1BaseResponse: + try: + buffer = BytesIO(await file.read()) + buffer.seek(0) + except Exception as e: + error_description = f"Не удалось сохранить файл для распознавания: {file.filename}, размер файла: {file.size}, по причине: {e}" + logger.error(error_description) + return V1BaseResponse( + success=False, + error_description=error_description, + data={} + ) + else: + logger.info(f"Получен и сохранён файл {file.filename}") + try: + result_dict = await asyncio.to_thread(process_file, buffer, params) + return V1BaseResponse( + success=result_dict.get('success', True), + error_description=result_dict.get('error_description'), + data={ + "raw_data": result_dict.get('raw_data'), + "sentenced_data": result_dict.get('sentenced_data'), + "diarized_data": result_dict.get('diarized_data') + } + ) + except Exception as e: + error_description = f"Ошибка обработки в process_file - {e}" + logger.error(error_description) + return V1BaseResponse( + success=False, + error_description=str(error_description), + data={} + ) + finally: + await file.close() + del file diff --git a/routes/v1/post_by_url.py b/routes/v1/post_by_url.py new file mode 100644 index 0000000..f065370 --- /dev/null +++ b/routes/v1/post_by_url.py @@ -0,0 +1,46 @@ +import uuid +import asyncio +from fastapi import APIRouter +from utils.pre_start_init import posted_and_downloaded_audio +from utils.do_logging import logger +from utils.get_audio_file import getting_audiofile, open_default_audiofile +from models.fast_api_models import SyncASRRequest, V1BaseResponse +from Recognizer.engine.file_recognition import process_file + +router = APIRouter() + +@router.post("/post_one_step_req", response_model=V1BaseResponse) +async def post_v1(params: SyncASRRequest) -> V1BaseResponse: + post_id = uuid.uuid4() + if params.AudioFileUrl: + res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) + else: + res, error_description = await open_default_audiofile(post_id) + + if not res: + logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') + return V1BaseResponse( + success=False, + error_description=error_description, + data={} + ) + + try: + result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) + return V1BaseResponse( + success=result_dict.get('success', True), + error_description=result_dict.get('error_description'), + data={ + "raw_data": result_dict.get('raw_data'), + "sentenced_data": result_dict.get('sentenced_data'), + "diarized_data": result_dict.get('diarized_data') + } + ) + except Exception as e: + error_description = f"Ошибка обработки в process_file - {e}" + logger.error(error_description) + return V1BaseResponse( + success=False, + error_description=str(error_description), + data={} + ) diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_http_routes.py b/tests/test_http_routes.py new file mode 100644 index 0000000..c43c0d9 --- /dev/null +++ b/tests/test_http_routes.py @@ -0,0 +1,105 @@ +import asyncio +import json + +import httpx + +from models.fast_api_models import V1BaseResponse as BaseResponse + +BASE_URL = "http://127.0.0.1:49153/v1" + + +async def _assert_base_response(body: dict, expect_success: bool): + assert "success" in body, f"Ответ должен содержать поле 'success' (V1BaseResponse)" + assert body["success"] is expect_success + + +async def test_post_by_url_success(): + async with httpx.AsyncClient() as client: + payload = { + "AudioFileUrl": "https://cdn.chatwm.opensmodel.sberdevices.ru/GigaAM/example.wav", + "keep_raw": True, + "do_echo_clearing": False, + "do_dialogue": False, + "do_punctuation": False, + } + resp = await client.post( + f"{BASE_URL}/post_one_step_req", json=payload, timeout=120.0 + ) + print("[POST /v1/post_one_step_req] status:", resp.status_code) + body = resp.json() + print(json.dumps(body, indent=2, ensure_ascii=False)) + + await _assert_base_response(body, expect_success=True) + return body + + +async def test_post_by_url_validation_error(): + """Ожидаем ErrorResponse при 422.""" + async with httpx.AsyncClient() as client: + payload = {"keep_raw": True} # отсутствует AudioFileUrl + resp = await client.post( + f"{BASE_URL}/post_one_step_req", json=payload, timeout=10.0 + ) + print("[POST /v1/post_one_step_req 422] status:", resp.status_code) + body = resp.json() + print(json.dumps(body, indent=2, ensure_ascii=False)) + + assert resp.status_code == 422 + await _assert_base_response(body, expect_success=False) + assert body.get("error_description") is not None + return body + + +async def test_post_by_file_success(): + async with httpx.AsyncClient() as client: + with open("../examples/orig.wav", "rb") as f: + files = {"file": ("orig.wav", f, "audio/wav")} + data = { + "keep_raw": "true", + "do_echo_clearing": "false", + "do_dialogue": "false", + "do_punctuation": "false", + "do_diarization": "false", + "diar_vad_sensity": "3", + } + resp = await client.post( + f"{BASE_URL}/post_file", data=data, files=files, timeout=120.0 + ) + print("[POST /v1/post_file] status:", resp.status_code) + body = resp.json() + print(json.dumps(body, indent=2, ensure_ascii=False)) + + await _assert_base_response(body, expect_success=True) + return body + + +async def test_post_by_file_validation_error(): + """Ожидаем ErrorResponse при 422 (нет файла).""" + async with httpx.AsyncClient() as client: + data = {"keep_raw": "true"} + resp = await client.post( + f"{BASE_URL}/post_file", data=data, timeout=10.0 + ) + print("[POST /v1/post_file 422] status:", resp.status_code) + body = resp.json() + print(json.dumps(body, indent=2, ensure_ascii=False)) + + assert resp.status_code == 422 + await _assert_base_response(body, expect_success=False) + return body + + +async def main(): + print("=== Запуск HTTP-тестов роутов ===\n") + await test_post_by_url_success() + print() + await test_post_by_url_validation_error() + print() + await test_post_by_file_success() + print() + await test_post_by_file_validation_error() + print("\n=== Все тесты завершены ===") + + +if __name__ == "__main__": + asyncio.run(main()) From 70efef270d9adba832b87c047c4d4fa947b5316b Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Tue, 28 Apr 2026 18:33:54 +0300 Subject: [PATCH 11/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20v1=20API=20=D0=B4=D0=BB=D1=8F=20?= =?UTF-8?q?=D1=81=D0=BE=D1=85=D1=80=D0=B0=D0=BD=D0=B5=D0=BD=D0=B8=D1=8F=20?= =?UTF-8?q?=D1=84=D1=83=D0=BD=D0=BA=D1=86=D0=B8=D0=BE=D0=BD=D0=B0=D0=BB?= =?UTF-8?q?=D0=B0=20=D0=B8=20=D0=BF=D1=80=D0=B8=D0=B2=D0=B5=D0=B4=D0=B5?= =?UTF-8?q?=D0=BD=D0=B8=D1=8F=20=D0=BE=D1=82=D0=B2=D0=B5=D1=82=D0=BE=D0=B2?= =?UTF-8?q?=20=D0=BA=20=D0=B5=D0=B4=D0=B8=D0=BD=D0=BE=D0=B9=20=D1=81=D1=82?= =?UTF-8?q?=D1=80=D1=83=D0=BA=D1=82=D1=83=D1=80=D0=B5.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 3 +-- main.py | 16 ++++++------ models/fast_api_models.py | 46 ++++++++++++++++++++++++++++++++++ routes/root.py | 24 +++++++++++------- routes/v1/is_alive.py | 34 +++++++++++++------------ routes/v1/post_by_file_FORM.py | 26 +++++++++---------- routes/v1/post_by_url.py | 26 +++++++++---------- routes/v1/post_ws.py | 10 ++++++++ routes/v1/root.py | 28 +++++++++++++++++++++ 9 files changed, 152 insertions(+), 61 deletions(-) create mode 100644 routes/v1/post_ws.py create mode 100644 routes/v1/root.py diff --git a/.gitignore b/.gitignore index e33dccd..c1eb734 100644 --- a/.gitignore +++ b/.gitignore @@ -19,5 +19,4 @@ /.idea/ .aider* /.continue/ -.env -/global_to_do.txt +.env \ No newline at end of file diff --git a/main.py b/main.py index 73c0bcc..1c9328c 100644 --- a/main.py +++ b/main.py @@ -114,14 +114,14 @@ async def general_exception_handler(request: Request, exc: Exception): app.mount("/static", StaticFiles(directory="static"), name="static") # Routers -app.include_router(root_router) -app.include_router(is_alive_router) -app.include_router(post_ws_router) -app.include_router(post_by_file_router) -app.include_router(post_by_url_router) -app.include_router(ws_audio_transkrib_router) -app.include_router(demo_router) -app.include_router(v1_router) +app.include_router(root_router, tags=["legacy"]) +app.include_router(is_alive_router, tags=["legacy"]) +app.include_router(post_ws_router, tags=["legacy"]) +app.include_router(post_by_file_router, tags=["legacy"]) +app.include_router(post_by_url_router, tags=["legacy"]) +app.include_router(ws_audio_transkrib_router, tags=["legacy"]) +app.include_router(demo_router, tags=["legacy"]) +app.include_router(v1_router, tags=["v1"]) def custom_openapi(): openapi_schema = get_openapi( diff --git a/models/fast_api_models.py b/models/fast_api_models.py index f52cc0e..e2f5e56 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -26,6 +26,52 @@ class V1BaseResponse(BaseModel): data: dict = {} +class RawData(BaseModel): + """Структура сырых данных ASR.""" + result: Optional[List[Dict[str, Any]]] = None + text: Optional[str] = None + + +class SentencedData(BaseModel): + """Структура разбитого на предложения ответа.""" + raw_text_sentenced_recognition: Optional[str] = None + list_of_sentenced_recognitions: Optional[List[Dict[str, Any]]] = None + full_text_only: Optional[List[str]] = None + err_state: Optional[Any] = None + + +class DiarizedData(BaseModel): + """Структура данных диаризации.""" + speakers: Optional[List[str]] = None + segments: Optional[List[Dict[str, Any]]] = None + + +class ASRData(BaseModel): + """Данные ответа ASR роутеров (post_by_url, post_by_file).""" + raw_data: Optional[RawData] = None + sentenced_data: Optional[SentencedData] = None + diarized_data: Optional[DiarizedData] = None + + +class V1ASRResponse(V1BaseResponse): + """Модель ответа для ASR роутов (post_by_url, post_by_file).""" + data: ASRData = ASRData() + + +class IsAliveData(BaseModel): + """Структура данных ответа is_alive.""" + state: str + tasks_in_work: int + free_memory_mb: Optional[float] = None + gpu_load_percent: Optional[float] = None + temperature_celsius: Optional[float] = None + + +class V1IsAliveResponse(V1BaseResponse): + """Модель ответа для is_alive.""" + data: Optional[IsAliveData] = None + + class ErrorResponse(BaseModel): """ Модель ошибки API. diff --git a/routes/root.py b/routes/root.py index 32ec78f..ca8a874 100644 --- a/routes/root.py +++ b/routes/root.py @@ -1,22 +1,28 @@ from fastapi import APIRouter -from models.fast_api_models import BaseResponse -import config - +from models.fast_api_models import V1BaseResponse +from config import settings router = APIRouter() -@router.get("/", response_model=BaseResponse) + +@router.get("/", response_model=V1BaseResponse) async def root(): - return BaseResponse( + """ + Корневой эндпоинт API v1. + + Returns: + V1BaseResponse: базовый ответ с приветственным сообщением. + """ + return V1BaseResponse( success=True, error_description=None, data={"message": "No_service_selected", "available_endpoints": { - "POST /post_one_step_req": "ASR by URL", - "POST /post_file": "ASR by file upload", + "POST v1/post_one_step_req": "ASR by URL", + "POST v1/post_file": "ASR by file upload", + "GET v1/is_alive": "Service health check", "WS /ws": "WebSocket streaming ASR", - "GET /is_alive": "Service health check", "GET /docs": "API documentation", "/demo": "DEMO UI page" }, - "try_addr": f"http://{config.settings.HOST}:{config.settings.PORT}/docs"} + "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs"} ) diff --git a/routes/v1/is_alive.py b/routes/v1/is_alive.py index a4db1d8..b2d2600 100644 --- a/routes/v1/is_alive.py +++ b/routes/v1/is_alive.py @@ -2,11 +2,11 @@ import logging from utils.pre_start_init import audio_to_asr from routes.is_alive import get_gpu_free_memory -from models.fast_api_models import V1BaseResponse +from models.fast_api_models import V1IsAliveResponse, IsAliveData router = APIRouter() -@router.get("/is_alive", response_model=V1BaseResponse) +@router.get("/is_alive", response_model=V1IsAliveResponse) async def check_if_service_is_alive_v1(): logging.info('GET /v1/is_alive') error_description = None @@ -16,18 +16,20 @@ async def check_if_service_is_alive_v1(): if error: error_description = error.get("error",None) - data = None + return V1IsAliveResponse( + success=True, + error_description=error_description, + data=None + ) else: - data = { - "state": state, - "tasks_in_work": tasks_in_work, - "free_memory_mb": free_mb, - "gpu_load_percent": gpu_load, - "temperature_celsius": temperature - } - - return V1BaseResponse( - success=True, - error_description=error_description, - data=data - ) + return V1IsAliveResponse( + success=True, + error_description=None, + data=IsAliveData( + state=state, + tasks_in_work=tasks_in_work, + free_memory_mb=free_mb, + gpu_load_percent=gpu_load, + temperature_celsius=temperature + ) + ) diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 449558e..5244e4a 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -3,48 +3,48 @@ from fastapi import APIRouter, Depends, File, Form, UploadFile from config import settings from utils.do_logging import logger -from models.fast_api_models import PostFileRequest, V1BaseResponse +from models.fast_api_models import PostFileRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file from routes.post_by_file_FORM import get_file_request router = APIRouter() -@router.post("/post_file", response_model=V1BaseResponse) +@router.post("/post_file", response_model=V1ASRResponse) async def async_receive_file_v1( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), -) -> V1BaseResponse: +) -> V1ASRResponse: try: buffer = BytesIO(await file.read()) buffer.seek(0) except Exception as e: error_description = f"Не удалось сохранить файл для распознавания: {file.filename}, размер файла: {file.size}, по причине: {e}" logger.error(error_description) - return V1BaseResponse( + return V1ASRResponse( success=False, error_description=error_description, - data={} + data=ASRData() ) else: logger.info(f"Получен и сохранён файл {file.filename}") try: result_dict = await asyncio.to_thread(process_file, buffer, params) - return V1BaseResponse( + return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), - data={ - "raw_data": result_dict.get('raw_data'), - "sentenced_data": result_dict.get('sentenced_data'), - "diarized_data": result_dict.get('diarized_data') - } + data=ASRData( + raw_data=RawData(**result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, + sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, + diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None + ) ) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - return V1BaseResponse( + return V1ASRResponse( success=False, error_description=str(error_description), - data={} + data=ASRData() ) finally: await file.close() diff --git a/routes/v1/post_by_url.py b/routes/v1/post_by_url.py index f065370..6293e6c 100644 --- a/routes/v1/post_by_url.py +++ b/routes/v1/post_by_url.py @@ -4,13 +4,13 @@ from utils.pre_start_init import posted_and_downloaded_audio from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile -from models.fast_api_models import SyncASRRequest, V1BaseResponse +from models.fast_api_models import SyncASRRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file router = APIRouter() -@router.post("/post_one_step_req", response_model=V1BaseResponse) -async def post_v1(params: SyncASRRequest) -> V1BaseResponse: +@router.post("/post_one_step_req", response_model=V1ASRResponse) +async def post_v1(params: SyncASRRequest) -> V1ASRResponse: post_id = uuid.uuid4() if params.AudioFileUrl: res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) @@ -19,28 +19,28 @@ async def post_v1(params: SyncASRRequest) -> V1BaseResponse: if not res: logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') - return V1BaseResponse( + return V1ASRResponse( success=False, error_description=error_description, - data={} + data=ASRData() ) try: result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) - return V1BaseResponse( + return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), - data={ - "raw_data": result_dict.get('raw_data'), - "sentenced_data": result_dict.get('sentenced_data'), - "diarized_data": result_dict.get('diarized_data') - } + data=ASRData( + raw_data=RawData(**result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, + sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, + diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None + ) ) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - return V1BaseResponse( + return V1ASRResponse( success=False, error_description=str(error_description), - data={} + data=ASRData() ) diff --git a/routes/v1/post_ws.py b/routes/v1/post_ws.py new file mode 100644 index 0000000..cd0d9ee --- /dev/null +++ b/routes/v1/post_ws.py @@ -0,0 +1,10 @@ +from fastapi import APIRouter, WebSocket +from fastapi.responses import JSONResponse +from models.fast_api_models import V1BaseResponse +import logging + +router = APIRouter() +logger = logging.getLogger(__name__) + +# TODO: Реализовать WebSocket-роут для потокового ASR +# Заглушка для будущей доработки в рамках Этапа 6 diff --git a/routes/v1/root.py b/routes/v1/root.py new file mode 100644 index 0000000..ca8a874 --- /dev/null +++ b/routes/v1/root.py @@ -0,0 +1,28 @@ +from fastapi import APIRouter +from models.fast_api_models import V1BaseResponse +from config import settings +router = APIRouter() + + +@router.get("/", response_model=V1BaseResponse) +async def root(): + """ + Корневой эндпоинт API v1. + + Returns: + V1BaseResponse: базовый ответ с приветственным сообщением. + """ + return V1BaseResponse( + success=True, + error_description=None, + data={"message": "No_service_selected", + "available_endpoints": { + "POST v1/post_one_step_req": "ASR by URL", + "POST v1/post_file": "ASR by file upload", + "GET v1/is_alive": "Service health check", + "WS /ws": "WebSocket streaming ASR", + "GET /docs": "API documentation", + "/demo": "DEMO UI page" + }, + "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs"} + ) From 81d7866a7c97810b67c27e71428a9c6c07107631 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Tue, 28 Apr 2026 19:22:41 +0300 Subject: [PATCH 12/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20legacy=20API=20=D0=B4=D0=BB=D1=8F=20?= =?UTF-8?q?=D1=81=D0=BE=D1=85=D1=80=D0=B0=D0=BD=D0=B5=D0=BD=D0=B8=D1=8F=20?= =?UTF-8?q?=D1=84=D1=83=D0=BD=D0=BA=D1=86=D0=B8=D0=BE=D0=BD=D0=B0=D0=BB?= =?UTF-8?q?=D0=B0=20=D0=B8=20=D0=BF=D1=80=D0=B8=D0=B2=D0=B5=D0=B4=D0=B5?= =?UTF-8?q?=D0=BD=D0=B8=D1=8F=20=D0=BE=D1=82=D0=B2=D0=B5=D1=82=D0=BE=D0=B2?= =?UTF-8?q?=20=D0=BA=20=D0=B5=D0=B4=D0=B8=D0=BD=D0=BE=D0=B9=20=D1=81=D1=82?= =?UTF-8?q?=D1=80=D1=83=D0=BA=D1=82=D1=83=D1=80=D0=B5.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 30 +++++---------------- models/fast_api_models.py | 1 + routes/__init__.py | 6 ----- routes/legacy/__init__.py | 14 ++++++++++ routes/{ => legacy}/demo_page.py | 0 routes/{ => legacy}/is_alive.py | 0 routes/{ => legacy}/post_by_file_FORM.py | 0 routes/{ => legacy}/post_by_url.py | 0 routes/{ => legacy}/post_ws.py | 0 routes/{ => legacy}/root.py | 0 routes/v1/__init__.py | 2 ++ routes/v1/is_alive.py | 2 +- routes/v1/post_by_file_FORM.py | 33 +++++++++++++++++++++--- 13 files changed, 55 insertions(+), 33 deletions(-) create mode 100644 routes/legacy/__init__.py rename routes/{ => legacy}/demo_page.py (100%) rename routes/{ => legacy}/is_alive.py (100%) rename routes/{ => legacy}/post_by_file_FORM.py (100%) rename routes/{ => legacy}/post_by_url.py (100%) rename routes/{ => legacy}/post_ws.py (100%) rename routes/{ => legacy}/root.py (100%) diff --git a/main.py b/main.py index 1c9328c..9977f83 100644 --- a/main.py +++ b/main.py @@ -15,14 +15,9 @@ from utils.pre_start_init import paths import threading -from routes.root import router as root_router -from routes.is_alive import router as is_alive_router -from routes.post_ws import router as post_ws_router -from routes.post_by_file_FORM import router as post_by_file_router -from routes.post_by_url import router as post_by_url_router from routes.ws_audio_transkrib import router as ws_audio_transkrib_router -from routes.demo_page import router as demo_router from routes.v1 import router as v1_router +from routes.legacy import router as legacy_router from models.fast_api_models import ErrorResponse import models @@ -64,10 +59,8 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE status_code=422, content=ErrorResponse( success=False, - error_description=str(exc), - raw_data={}, - sentenced_data={}, - diarized_data={}, + error_description="Validation error", + details=str(exc), ).model_dump() ) @@ -79,24 +72,20 @@ async def http_exception_handler(request: Request, exc: StarletteHTTPException): content=ErrorResponse( success=False, error_description=exc.detail, - raw_data={}, - sentenced_data={}, - diarized_data={}, ).model_dump() ) @app.exception_handler(Exception) async def general_exception_handler(request: Request, exc: Exception): - logger.error(f"Unhandled exception: {exc}") + import traceback + logger.error(f"Unhandled exception: {exc}\n{traceback.format_exc()}") return JSONResponse( status_code=500, content=ErrorResponse( success=False, error_description="Internal server error", - raw_data={}, - sentenced_data={}, - diarized_data={}, + details=str(exc), ).model_dump() ) @@ -114,13 +103,8 @@ async def general_exception_handler(request: Request, exc: Exception): app.mount("/static", StaticFiles(directory="static"), name="static") # Routers -app.include_router(root_router, tags=["legacy"]) -app.include_router(is_alive_router, tags=["legacy"]) -app.include_router(post_ws_router, tags=["legacy"]) -app.include_router(post_by_file_router, tags=["legacy"]) -app.include_router(post_by_url_router, tags=["legacy"]) app.include_router(ws_audio_transkrib_router, tags=["legacy"]) -app.include_router(demo_router, tags=["legacy"]) +app.include_router(legacy_router, tags=["legacy"]) app.include_router(v1_router, tags=["v1"]) def custom_openapi(): diff --git a/models/fast_api_models.py b/models/fast_api_models.py index e2f5e56..bd40395 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -78,6 +78,7 @@ class ErrorResponse(BaseModel): """ success: bool = False error_description: str + details: Optional[str] = None raw_data: Optional[Dict[str, Any]] = None sentenced_data: Optional[Dict[str, Any]] = None diarized_data: Optional[Dict[str, Any]] = None diff --git a/routes/__init__.py b/routes/__init__.py index 00627e3..934f058 100644 --- a/routes/__init__.py +++ b/routes/__init__.py @@ -1,7 +1 @@ -from . import root -from . import is_alive from . import ws_audio_transkrib -from . import post_ws -from . import post_by_url -from . import post_by_file_FORM -from . import demo_page diff --git a/routes/legacy/__init__.py b/routes/legacy/__init__.py new file mode 100644 index 0000000..9f72e43 --- /dev/null +++ b/routes/legacy/__init__.py @@ -0,0 +1,14 @@ +from fastapi import APIRouter +from .is_alive import router as is_alive_router +from .post_by_url import router as post_by_url_router +from .post_by_file_FORM import router as post_by_file_router +from .root import router as root_router +from .demo_page import router as demo_page_router + +router = APIRouter() +router.include_router(is_alive_router) +router.include_router(post_by_url_router) +router.include_router(post_by_file_router) +router.include_router(post_by_file_router) +router.include_router(root_router) +router.include_router(demo_page_router) diff --git a/routes/demo_page.py b/routes/legacy/demo_page.py similarity index 100% rename from routes/demo_page.py rename to routes/legacy/demo_page.py diff --git a/routes/is_alive.py b/routes/legacy/is_alive.py similarity index 100% rename from routes/is_alive.py rename to routes/legacy/is_alive.py diff --git a/routes/post_by_file_FORM.py b/routes/legacy/post_by_file_FORM.py similarity index 100% rename from routes/post_by_file_FORM.py rename to routes/legacy/post_by_file_FORM.py diff --git a/routes/post_by_url.py b/routes/legacy/post_by_url.py similarity index 100% rename from routes/post_by_url.py rename to routes/legacy/post_by_url.py diff --git a/routes/post_ws.py b/routes/legacy/post_ws.py similarity index 100% rename from routes/post_ws.py rename to routes/legacy/post_ws.py diff --git a/routes/root.py b/routes/legacy/root.py similarity index 100% rename from routes/root.py rename to routes/legacy/root.py diff --git a/routes/v1/__init__.py b/routes/v1/__init__.py index 0fdb9b8..d26a7c1 100644 --- a/routes/v1/__init__.py +++ b/routes/v1/__init__.py @@ -2,8 +2,10 @@ from .is_alive import router as is_alive_router from .post_by_url import router as post_by_url_router from .post_by_file_FORM import router as post_by_file_router +from .root import router as root_router router = APIRouter(prefix="/v1") router.include_router(is_alive_router) router.include_router(post_by_url_router) router.include_router(post_by_file_router) +router.include_router(root_router) diff --git a/routes/v1/is_alive.py b/routes/v1/is_alive.py index b2d2600..9b17702 100644 --- a/routes/v1/is_alive.py +++ b/routes/v1/is_alive.py @@ -1,7 +1,7 @@ from fastapi import APIRouter import logging from utils.pre_start_init import audio_to_asr -from routes.is_alive import get_gpu_free_memory +from routes.legacy.is_alive import get_gpu_free_memory from models.fast_api_models import V1IsAliveResponse, IsAliveData router = APIRouter() diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 5244e4a..0993636 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -1,13 +1,40 @@ from io import BytesIO import asyncio -from fastapi import APIRouter, Depends, File, Form, UploadFile -from config import settings from utils.do_logging import logger +from config import settings +from fastapi import APIRouter, Depends, File, Form, UploadFile from models.fast_api_models import PostFileRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file -from routes.post_by_file_FORM import get_file_request router = APIRouter() +# Функция для извлечения параметров из FormData +def get_file_request( + keep_raw: bool = Form(default=True, description="Сохранять сырые данные."), + do_echo_clearing: bool = Form(default=False, description="Убирать межканальное эхо."), + do_dialogue: bool = Form(default=False, description="Строить диалог."), + do_punctuation: bool = Form(default=False, description="Восстанавливать пунктуацию."), + do_diarization: bool = Form(default=False, description="Разделять по спикерам."), + diar_vad_sensity: int = Form(default=3, description="Чувствительность VAD."), + use_batch: bool = Form(default=settings.USE_BATCH, description="Использовать батчинг для ASR."), + batch_size: int = Form(default=settings.ASR_BATCH_SIZE, description="Размер батча для ASR."), + do_auto_speech_speed_correction: bool = Form(default=True, description="Корректировать скорость речи при распознавании."), + speech_speed_correction_multiplier: float = Form(default=1, description="Базовый коэффициент скорости речи."), + make_mono: bool = Form(default=False, description="Соединить несколько каналов в mono"), +) -> PostFileRequest: + return PostFileRequest( + keep_raw=keep_raw, + do_echo_clearing=do_echo_clearing, + do_dialogue=do_dialogue, + do_punctuation=do_punctuation, + do_diarization=do_diarization, + use_batch=use_batch, + diar_vad_sensity=diar_vad_sensity, + batch_size=batch_size, + make_mono=make_mono, + do_auto_speech_speed_correction = do_auto_speech_speed_correction, + speech_speed_correction_multiplier = speech_speed_correction_multiplier + ) + @router.post("/post_file", response_model=V1ASRResponse) async def async_receive_file_v1( From d0c79c9743b0fb4a329a2ac1b5a480fbfcab2df7 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 29 Apr 2026 17:29:36 +0300 Subject: [PATCH 13/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D1=8F=20=D0=BE=D0=BF=D0=B8=D1=81=D0=B0=D0=BD?= =?UTF-8?q?=D0=B8=D1=8F=20"/ws"?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.py | 125 ++++++++++++++++++++--------- examples/streaming_client.py | 5 +- main.py | 69 ++++++++-------- routes/legacy/__init__.py | 5 +- routes/legacy/is_alive.py | 3 - routes/legacy/post_by_file_FORM.py | 2 +- routes/legacy/post_ws.py | 9 --- routes/legacy/root.py | 9 +-- routes/v1/__init__.py | 2 +- routes/v1/post_by_file_FORM.py | 2 +- routes/v1/post_ws.py | 10 --- 11 files changed, 134 insertions(+), 107 deletions(-) delete mode 100644 routes/legacy/post_ws.py delete mode 100644 routes/v1/post_ws.py diff --git a/config.py b/config.py index 817d735..a8ddec0 100644 --- a/config.py +++ b/config.py @@ -145,42 +145,7 @@ def _compute_derived(self): # Единственный экземпляр настроек settings = Settings() -# ============================================================================== -# Обратная совместимость: экспортируем атрибуты на уровень модуля. -# Все остальные модули проекта могут продолжать использовать `config.HOST`, -# `config.PORT` и т.д. без изменений. -# ============================================================================== -# BASE_SAMPLE_RATE = settings.BASE_SAMPLE_RATE -# PROVIDER = settings.PROVIDER -# NUM_THREADS = settings.NUM_THREADS -# HF_HOME = settings.HF_HOME -# LOGGING_LEVEL = settings.LOGGING_LEVEL -# LOGGING_FORMAT = settings.LOGGING_FORMAT -# FILENAME = settings.FILENAME -# FILEMODE = settings.FILEMODE -# LOG_BACKUP_COUNT = settings.LOG_BACKUP_COUNT -# IS_PROD = settings.IS_PROD -# MAX_OVERLAP_DURATION = settings.MAX_OVERLAP_DURATION -# RECOGNITION_ATTEMPTS = settings.RECOGNITION_ATTEMPTS -# SPEECH_PER_SEC_NORM_RATE = settings.SPEECH_PER_SEC_NORM_RATE -# MAKE_MONO = settings.MAKE_MONO -# USE_BATCH = settings.USE_BATCH -# ASR_BATCH_SIZE = settings.ASR_BATCH_SIZE -# VAD_SENSITIVITY = settings.VAD_SENSITIVITY -# VAD_WITH_GPU = settings.VAD_WITH_GPU -# BETWEEN_WORDS_PERCENTILE = settings.BETWEEN_WORDS_PERCENTILE -# CAN_PUNCTUATE = settings.CAN_PUNCTUATE -# PUNCTUATE_WITH_GPU = settings.PUNCTUATE_WITH_GPU -# CAN_DIAR = settings.CAN_DIAR -# DIAR_MODEL_NAME = settings.DIAR_MODEL_NAME -# DIAR_WITH_GPU = settings.DIAR_WITH_GPU -# CPU_WORKERS = settings.CPU_WORKERS -# DIAR_GPU_BATCH_SIZE = settings.DIAR_GPU_BATCH_SIZE -# DO_SPEED_SPEECH_CORRECTION = settings.DO_SPEED_SPEECH_CORRECTION -# SPEED_SPEECH_CORRECTION_MULTIPLIER = settings.SPEED_SPEECH_CORRECTION_MULTIPLIER -# DO_LOCAL_FILE_RECOGNITIONS = settings.DO_LOCAL_FILE_RECOGNITIONS -# DELETE_LOCAL_FILE_AFTR_ASR = settings.DELETE_LOCAL_FILE_AFTR_ASR -# HUMAN_FORMAT_MD_FILE = settings.HUMAN_FORMAT_MD_FILE + AUDIOEXTENTIONS = [ # Основные форматы @@ -196,3 +161,91 @@ def _compute_derived(self): # Редкие/устаревшие форматы '.669', '.mtm', '.med', '.far', '.umx' ] + +# Описание WebSocket для OpenAPI +WS_DESCRIPTION = """ +## WebSocket Endpoint - `/ws` +### Пример конфигурации + +Отправьте JSON с конфигурацией: + +```json + { + "config": { + "audio_format": "pcm16", + "sample_rate": 16000, + "wait_null_answers": true, + "do_dialogue": false, + "do_punctuation": false, + "channelName": "channel_1" # id канала из астериск, например. + } + } +``` + +### Пример передачи данных. + +Периодически отправляйте raw_audio_data - PCM, 16-bit, mono + +```json + { + "bytes": binary + } +``` + +### Пример EOF + +По завершении отправьте: +```json + { + "text": "eof" + } +``` + +### Ответы + +``` json + { + 'channel_name': 'Null', + 'silence': False, + 'data': { + 'result': + [ + {'conf': 1.0, 'start': 116.48, 'end': 116.76, 'word': 'владимир'}, + {'conf': 1.0, 'start': 116.92, 'end': 117.48, 'word': 'анатольевич'} + ], + 'text': 'владимир анатольевич'}, + 'error': None, + 'last_message': False, + 'sentenced_data': {} + } +``` +Если в config передать "do_dialogue":true и "do_punctuation":true то в последнем ответе будет предоставлен Капитализированный +текст, с пунктуацией разбитый на фразы. + +```json + { + 'channel_name': 'Null', + 'silence': False, + 'data': { + 'result': + [ + {"conf": 1, "start": 0.04, "end": 0.36, "word": "ничьих"}, + {"conf": 1, "start": 0.52, "end": 0.56, "word": "не"}, + {"conf": 1, "start": 0.64, "end": 0.92, "word": "требуя" }, + {"conf": 1, "start": 1.08,"end": 1.44,"word": "похвал"}, + ], + "text": "ничьих не требуя похвал ... " + }, + 'error': None, + 'last_message': True, + 'sentenced_data': { + 'raw_text_sentenced_recognition': "channel_1: Ничьих, не требуя ... мои.\n channel_1: У Лукоморья дуб зеленый.", # текст построчно разбитый на фразы. + 'list_of_sentenced_recognitions': [{'start': 1.0, 'end': 1.28, 'text': 'У Лукоморья дуб зеленый.', 'speaker': 'channel_1'},... ] + "full_text_only": [ + "Ничьих, не требуя похвал. Счастлив уж я надеждой сладкой, что дева с трепетом любви посмотрит, может быть, украдкой на песни грешные мои. У Лукоморья дуб зеленый." + ], + } + } + +``` +""" \ No newline at end of file diff --git a/examples/streaming_client.py b/examples/streaming_client.py index 72a6adb..47356be 100644 --- a/examples/streaming_client.py +++ b/examples/streaming_client.py @@ -47,7 +47,9 @@ async def _send_config(self, websocket, sample_rate, wait_null_answers): "sample_rate": sample_rate, "wait_null_answers": wait_null_answers, "audio_format": "pcm16", - "language": "ru" + "language": "ru", + "do_dialogue": "true", + "do_punctuation": "false", } await websocket.send(ujson.dumps({"config": config})) logging.info("Configuration sent") @@ -136,6 +138,7 @@ async def _handle_responses(self, websocket): args.uri, frame_rate=args.frame_rate, buffer_size_sec=args.buffer_size + ) try: diff --git a/main.py b/main.py index 9977f83..dedcfab 100644 --- a/main.py +++ b/main.py @@ -20,7 +20,7 @@ from routes.legacy import router as legacy_router from models.fast_api_models import ErrorResponse import models - +from config import WS_DESCRIPTION @asynccontextmanager async def lifespan(app): @@ -49,9 +49,9 @@ async def lifespan(app): version="1.0", docs_url='/docs', root_path='/root', - title='ASR' -) - + title='ASR', + description=WS_DESCRIPTION + ) @app.exception_handler(RequestValidationError) async def validation_exception_handler(request: Request, exc: RequestValidationError): @@ -107,44 +107,41 @@ async def general_exception_handler(request: Request, exc: Exception): app.include_router(legacy_router, tags=["legacy"]) app.include_router(v1_router, tags=["v1"]) -def custom_openapi(): - openapi_schema = get_openapi( - title="ASR Speech Recognition API", - version="1.0.0", - description="Real-time Russian ASR via WebSocket. Send raw audio chunks (16 kHz, mono).", - routes=app.routes, - contact={"email": "Kojevnikov@amulex.ru"}, - ) - # Добавляем пример для WebSocket - openapi_schema["paths"]["/ws"]["websocket"] = { - "summary": "Stream audio for transcription", - "requestBody": { - "content": { - "audio/*": { - "example": {"description": "Raw audio bytes (PCM, 16-bit)"} - } - } - }, - "responses": { - "200": { - "description": "ASR result in JSON", - "content": { - "application/json": { - "example": {"text": "привет мир", "confidence": 0.95} - } - } - } - } - } - app.openapi_schema = openapi_schema - return openapi_schema +# def custom_openapi(): +# if app.openapi_schema: +# return app.openapi_schema +# openapi_schema = get_openapi( +# title="ASR Speech Recognition API", +# version="1.0.0", +# description="Real-time Russian ASR via WebSocket. Send raw audio chunks (16 kHz, mono).", +# routes=app.routes, +# contact={"email": "Kojevnikov@amulex.ru"}, +# ) +# # Добавляем описание для WebSocket endpoint +# # FastAPI автоматически генерирует WebSocket документацию, но мы можем улучшить её описание +# if "/ws" in openapi_schema.get("paths", {}): +# # Удаляем стандартный путь, так как он не подходит для WebSocket +# del openapi_schema["paths"]["/ws"] +# +# # Добавляем кастомное описание WebSocket +# openapi_schema["paths"]["/ws"] = { +# "summary": "WebSocket Stream for Audio Transcription", +# "description": "Подключитесь по WebSocket для потоковой передачи аудио. Отправляйте raw audio bytes (PCM, 16-bit, mono, 16kHz). Получайте результаты транскрипции в JSON.", +# "servers": [ +# {"url": "ws://{host}/ws", "description": "WebSocket server"} +# ], +# "x-postman-collection-name": "ASR WebSocket" +# } +# +# app.openapi_schema = openapi_schema +# return openapi_schema try: if __name__ == '__main__': - app.openapi = custom_openapi + # app.openapi = app.openapi_schema uvicorn.run(app, host=settings.HOST, port=settings.PORT) except KeyboardInterrupt: logger.info('\nDone') diff --git a/routes/legacy/__init__.py b/routes/legacy/__init__.py index 9f72e43..4e13035 100644 --- a/routes/legacy/__init__.py +++ b/routes/legacy/__init__.py @@ -6,9 +6,8 @@ from .demo_page import router as demo_page_router router = APIRouter() +router.include_router(root_router) +router.include_router(demo_page_router) router.include_router(is_alive_router) router.include_router(post_by_url_router) router.include_router(post_by_file_router) -router.include_router(post_by_file_router) -router.include_router(root_router) -router.include_router(demo_page_router) diff --git a/routes/legacy/is_alive.py b/routes/legacy/is_alive.py index afec665..3f61eba 100644 --- a/routes/legacy/is_alive.py +++ b/routes/legacy/is_alive.py @@ -1,10 +1,7 @@ from fastapi import APIRouter import logging -import datetime -import os import pynvml from utils.pre_start_init import audio_to_asr -from models.fast_api_models import BaseResponse router = APIRouter() diff --git a/routes/legacy/post_by_file_FORM.py b/routes/legacy/post_by_file_FORM.py index 961f551..3d9143d 100644 --- a/routes/legacy/post_by_file_FORM.py +++ b/routes/legacy/post_by_file_FORM.py @@ -45,7 +45,7 @@ def get_file_request( @router.post("/post_file", response_model=BaseResponse) -async def async_receive_file( +async def async_receive_file_legacy( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), ) -> BaseResponse: diff --git a/routes/legacy/post_ws.py b/routes/legacy/post_ws.py deleted file mode 100644 index be09dfd..0000000 --- a/routes/legacy/post_ws.py +++ /dev/null @@ -1,9 +0,0 @@ -from fastapi import APIRouter -from models.fast_api_models import WebSocketModel - -router = APIRouter() - -@router.post("/ws") -async def post_not_websocket(ws:WebSocketModel): - """Описание для вебсокета ниже в описании WebSocketModel """ - return f"Прочти инструкцию в Schemas - 'WebSocketModel'" diff --git a/routes/legacy/root.py b/routes/legacy/root.py index ca8a874..a12b080 100644 --- a/routes/legacy/root.py +++ b/routes/legacy/root.py @@ -4,7 +4,7 @@ router = APIRouter() -@router.get("/", response_model=V1BaseResponse) +@router.get("/") async def root(): """ Корневой эндпоинт API v1. @@ -12,10 +12,7 @@ async def root(): Returns: V1BaseResponse: базовый ответ с приветственным сообщением. """ - return V1BaseResponse( - success=True, - error_description=None, - data={"message": "No_service_selected", + return {"message": "No_service_selected", "available_endpoints": { "POST v1/post_one_step_req": "ASR by URL", "POST v1/post_file": "ASR by file upload", @@ -25,4 +22,4 @@ async def root(): "/demo": "DEMO UI page" }, "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs"} - ) + diff --git a/routes/v1/__init__.py b/routes/v1/__init__.py index d26a7c1..d335a31 100644 --- a/routes/v1/__init__.py +++ b/routes/v1/__init__.py @@ -5,7 +5,7 @@ from .root import router as root_router router = APIRouter(prefix="/v1") +router.include_router(root_router) router.include_router(is_alive_router) router.include_router(post_by_url_router) router.include_router(post_by_file_router) -router.include_router(root_router) diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 0993636..9e6dc0e 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -37,7 +37,7 @@ def get_file_request( @router.post("/post_file", response_model=V1ASRResponse) -async def async_receive_file_v1( +async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), ) -> V1ASRResponse: diff --git a/routes/v1/post_ws.py b/routes/v1/post_ws.py deleted file mode 100644 index cd0d9ee..0000000 --- a/routes/v1/post_ws.py +++ /dev/null @@ -1,10 +0,0 @@ -from fastapi import APIRouter, WebSocket -from fastapi.responses import JSONResponse -from models.fast_api_models import V1BaseResponse -import logging - -router = APIRouter() -logger = logging.getLogger(__name__) - -# TODO: Реализовать WebSocket-роут для потокового ASR -# Заглушка для будущей доработки в рамках Этапа 6 From 431565e55a6206d90fbae0c0accaf016c9598177 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 29 Apr 2026 19:53:09 +0300 Subject: [PATCH 14/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20ALLOWED=5FHOSTS,=20CORS=5FORIGINS,?= =?UTF-8?q?=20TRUSTED=5FPROXIES.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 2 +- .idea/Vosk5_FastAPI_streaming.iml | 1 + config.py | 14 ++++- main.py | 71 +++++++++++++------------- requirements.txt | 3 ++ tests/test_cors_middleware.py | 36 +++++++++++++ tests/test_gzip_middleware.py | 30 +++++++++++ tests/test_proxy_headers_middleware.py | 22 ++++++++ tests/test_request_id_middleware.py | 23 +++++++++ tests/test_trusted_host_middleware.py | 29 +++++++++++ 10 files changed, 194 insertions(+), 37 deletions(-) create mode 100644 tests/test_cors_middleware.py create mode 100644 tests/test_gzip_middleware.py create mode 100644 tests/test_proxy_headers_middleware.py create mode 100644 tests/test_request_id_middleware.py create mode 100644 tests/test_trusted_host_middleware.py diff --git a/.gitignore b/.gitignore index c1eb734..ed6794d 100644 --- a/.gitignore +++ b/.gitignore @@ -4,7 +4,7 @@ *.ogg *.onnx *.m4a -*.md +*.history.md /.venv/ /venv/ /models/vosk-model-ru/ diff --git a/.idea/Vosk5_FastAPI_streaming.iml b/.idea/Vosk5_FastAPI_streaming.iml index db89304..7652601 100644 --- a/.idea/Vosk5_FastAPI_streaming.iml +++ b/.idea/Vosk5_FastAPI_streaming.iml @@ -10,6 +10,7 @@ + diff --git a/config.py b/config.py index a8ddec0..0bce7a4 100644 --- a/config.py +++ b/config.py @@ -1,3 +1,4 @@ +import logging import os from datetime import date from pydantic_settings import BaseSettings, SettingsConfigDict @@ -19,6 +20,9 @@ class Settings(BaseSettings): # server settings HOST: str = '0.0.0.0' PORT: int = 49153 + ALLOWED_HOSTS: list[str] = ["*"] + CORS_ORIGINS: list[str] = ["*"] + TRUSTED_PROXIES: list[str] = ["*"] # Model settings # Vosk5SmallStreaming Vosk5 Gigaam Whisper Gigaam_rnnt, "gigaam-v3-rnnt", "gigaam-v3-ctc" @@ -139,6 +143,14 @@ def _compute_derived(self): if self.PUNCTUATE_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: self.PUNCTUATE_WITH_GPU = False + # Предупреждение о небезопасной CORS-конфигурации в продакшене + if self.IS_PROD and self.CORS_ORIGINS == ["*"]: + logging.warning( + "SECURITY WARNING: CORS_ORIGINS is set to ['*'] in production (IS_PROD=True). " + "This is insecure when allow_credentials=True. " + "Please specify explicit origins in CORS_ORIGINS." + ) + return self @@ -248,4 +260,4 @@ def _compute_derived(self): } ``` -""" \ No newline at end of file +""" diff --git a/main.py b/main.py index dedcfab..562d66f 100644 --- a/main.py +++ b/main.py @@ -10,7 +10,11 @@ from starlette.exceptions import HTTPException as StarletteHTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles -from fastapi.openapi.utils import get_openapi +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.middleware.gzip import GZipMiddleware +from starlette.middleware.trustedhost import TrustedHostMiddleware +from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware +import uuid from utils.files_whatcher import start_file_watcher from utils.pre_start_init import paths import threading @@ -22,6 +26,15 @@ import models from config import WS_DESCRIPTION +class RequestIDMiddleware(BaseHTTPMiddleware): + async def dispatch(self, request: Request, call_next): + request_id = request.headers.get("X-Request-ID", str(uuid.uuid4())) + request.state.request_id = request_id + response = await call_next(request) + response.headers["X-Request-ID"] = request_id + return response + + @asynccontextmanager async def lifespan(app): # on_start @@ -90,15 +103,36 @@ async def general_exception_handler(request: Request, exc: Exception): ) +# RequestID middleware +app.add_middleware(RequestIDMiddleware) + +# ProxyHeaders middleware +app.add_middleware( + ProxyHeadersMiddleware, + trusted_hosts=settings.TRUSTED_PROXIES, +) + +# TrustedHost middleware +app.add_middleware( + TrustedHostMiddleware, + allowed_hosts=settings.ALLOWED_HOSTS, +) + # CORS middleware app.add_middleware( CORSMiddleware, - allow_origins=["*"], + allow_origins=settings.CORS_ORIGINS, allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) +# GZip middleware +app.add_middleware( + GZipMiddleware, + minimum_size=1000, +) + # Static files app.mount("/static", StaticFiles(directory="static"), name="static") @@ -107,38 +141,6 @@ async def general_exception_handler(request: Request, exc: Exception): app.include_router(legacy_router, tags=["legacy"]) app.include_router(v1_router, tags=["v1"]) -# def custom_openapi(): -# if app.openapi_schema: -# return app.openapi_schema -# openapi_schema = get_openapi( -# title="ASR Speech Recognition API", -# version="1.0.0", -# description="Real-time Russian ASR via WebSocket. Send raw audio chunks (16 kHz, mono).", -# routes=app.routes, -# contact={"email": "Kojevnikov@amulex.ru"}, -# ) -# # Добавляем описание для WebSocket endpoint -# # FastAPI автоматически генерирует WebSocket документацию, но мы можем улучшить её описание -# if "/ws" in openapi_schema.get("paths", {}): -# # Удаляем стандартный путь, так как он не подходит для WebSocket -# del openapi_schema["paths"]["/ws"] -# -# # Добавляем кастомное описание WebSocket -# openapi_schema["paths"]["/ws"] = { -# "summary": "WebSocket Stream for Audio Transcription", -# "description": "Подключитесь по WebSocket для потоковой передачи аудио. Отправляйте raw audio bytes (PCM, 16-bit, mono, 16kHz). Получайте результаты транскрипции в JSON.", -# "servers": [ -# {"url": "ws://{host}/ws", "description": "WebSocket server"} -# ], -# "x-postman-collection-name": "ASR WebSocket" -# } -# -# app.openapi_schema = openapi_schema -# return openapi_schema - - - - try: if __name__ == '__main__': # app.openapi = app.openapi_schema @@ -147,4 +149,3 @@ async def general_exception_handler(request: Request, exc: Exception): logger.info('\nDone') except Exception as e: logger.error(f'\nDone with error {e}') - diff --git a/requirements.txt b/requirements.txt index ed26c38..d59a477 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,4 @@ +pytest setuptools pydantic_settings httpx~=0.28.1 @@ -47,3 +48,5 @@ watchdog~=6.0.0 onnx-asr == 0.10.2 hf_xet onnxruntime + +starlette \ No newline at end of file diff --git a/tests/test_cors_middleware.py b/tests/test_cors_middleware.py new file mode 100644 index 0000000..091ff1f --- /dev/null +++ b/tests/test_cors_middleware.py @@ -0,0 +1,36 @@ +import pytest +from fastapi.testclient import TestClient +from main import app +from config import settings + + +@pytest.fixture +def client(): + with TestClient(app) as c: + yield c + + +def test_cors_origins_default_config(): + """По умолчанию CORS_ORIGINS должен быть ['*'].""" + assert hasattr(settings, "CORS_ORIGINS") + assert settings.CORS_ORIGINS == ["*"] + + +def test_cors_middleware_allows_any_origin(client): + """При CORS_ORIGINS=['*'] любой Origin должен получить разрешение.""" + response = client.get("/", headers={"Origin": "http://example.com"}) + assert "access-control-allow-origin" in response.headers + assert response.headers["access-control-allow-origin"] == "*" + + +def test_cors_preflight_request(client): + """Preflight OPTIONS запрос должен возвращать 200 с CORS-заголовками.""" + response = client.options( + "/", + headers={ + "Origin": "http://example.com", + "Access-Control-Request-Method": "GET", + }, + ) + assert response.status_code == 200 + assert "access-control-allow-origin" in response.headers diff --git a/tests/test_gzip_middleware.py b/tests/test_gzip_middleware.py new file mode 100644 index 0000000..3c2f872 --- /dev/null +++ b/tests/test_gzip_middleware.py @@ -0,0 +1,30 @@ +import pytest +from fastapi.testclient import TestClient +from starlette.middleware.gzip import GZipMiddleware +from main import app + + +@pytest.fixture +def client(): + with TestClient(app) as c: + yield c + + +def test_gzip_middleware_is_configured(): + """GZipMiddleware должна быть зарегистрирована в приложении.""" + middleware_classes = [m.cls for m in app.user_middleware] + assert GZipMiddleware in middleware_classes + + +def test_gzip_middleware_compresses_large_response(client): + """Большие ответы (>1000 байт) должны сжиматься при Accept-Encoding: gzip.""" + response = client.get("/openapi.json", headers={"Accept-Encoding": "gzip"}) + assert response.status_code == 200 + assert len(response.content) > 0 + + +def test_gzip_middleware_skips_small_response(client): + """Маленькие ответы (<1000 байт) не должны сжиматься.""" + response = client.get("/", headers={"Accept-Encoding": "gzip"}) + assert response.status_code == 200 + assert "content-encoding" not in response.headers diff --git a/tests/test_proxy_headers_middleware.py b/tests/test_proxy_headers_middleware.py new file mode 100644 index 0000000..0391719 --- /dev/null +++ b/tests/test_proxy_headers_middleware.py @@ -0,0 +1,22 @@ +import pytest +from fastapi.testclient import TestClient +from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware +from main import app + + +@pytest.fixture +def client(): + with TestClient(app) as c: + yield c + + +def test_proxy_headers_middleware_is_configured(): + """ProxyHeadersMiddleware должна быть зарегистрирована в приложении.""" + middleware_classes = [m.cls for m in app.user_middleware] + assert ProxyHeadersMiddleware in middleware_classes + + +def test_proxy_headers_middleware_allows_forwarded_header(client): + """При TRUSTED_PROXIES=['*'] запрос с X-Forwarded-For не должен падать.""" + response = client.get("/", headers={"X-Forwarded-For": "203.0.113.42"}) + assert response.status_code == 200 diff --git a/tests/test_request_id_middleware.py b/tests/test_request_id_middleware.py new file mode 100644 index 0000000..e561c50 --- /dev/null +++ b/tests/test_request_id_middleware.py @@ -0,0 +1,23 @@ +import pytest +from fastapi.testclient import TestClient +from main import app + + +@pytest.fixture +def client(): + with TestClient(app) as c: + yield c + + +def test_request_id_middleware_generates_id(client): + """Если клиент не передал X-Request-ID, middleware должен сгенерировать UUID.""" + response = client.get("/") + assert "x-request-id" in response.headers + assert len(response.headers["x-request-id"]) > 0 + + +def test_request_id_middleware_preserves_provided_id(client): + """Если клиент передал X-Request-ID, он должен вернуться в ответе без изменений.""" + custom_id = "my-custom-id-12345" + response = client.get("/", headers={"X-Request-ID": custom_id}) + assert response.headers["x-request-id"] == custom_id diff --git a/tests/test_trusted_host_middleware.py b/tests/test_trusted_host_middleware.py new file mode 100644 index 0000000..522a05d --- /dev/null +++ b/tests/test_trusted_host_middleware.py @@ -0,0 +1,29 @@ +import pytest +from fastapi.testclient import TestClient +from main import app +from config import settings + + +@pytest.fixture +def restricted_client(): + """Клиент с ограниченным ALLOWED_HOSTS для проверки TrustedHostMiddleware.""" + original_hosts = settings.ALLOWED_HOSTS.copy() + settings.ALLOWED_HOSTS[:] = ["trusted.example.com"] + # Сбрасываем кэш middleware stack, чтобы перестроить с актуальными настройками + app._middleware_stack = None + with TestClient(app) as client: + yield client + settings.ALLOWED_HOSTS[:] = original_hosts + app._middleware_stack = None + + +def test_trusted_host_middleware_rejects_invalid_host(restricted_client): + """Запросы с недоверенным Host должны отклоняться с 400.""" + response = restricted_client.get("/", headers={"Host": "untrusted.example.com"}) + assert response.status_code == 400 + + +def test_trusted_host_middleware_allows_valid_host(restricted_client): + """Запросы с доверенным Host должны проходить (не 400).""" + response = restricted_client.get("/", headers={"Host": "trusted.example.com"}) + assert response.status_code != 400 From a3b62f2cd90671703bd00af862444ff2144364af Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 30 Apr 2026 13:00:17 +0300 Subject: [PATCH 15/46] =?UTF-8?q?=D0=BF=D0=B5=D1=80=D0=B5=D0=B5=D0=B7?= =?UTF-8?q?=D0=B4=20=D0=BB=D0=BE=D0=B3=D0=B3=D0=B5=D1=80=D0=B0=20=D0=B8=20?= =?UTF-8?q?recognizer=20=D0=B2=20lifespan?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/__init__.py | 5 +- Recognizer/engine/file_recognition.py | 9 +- Recognizer/engine/stream_recognition.py | 6 +- core/__init__.py | 1 + core/logging_config.py | 97 ++++++++++ main.py | 34 +++- routes/legacy/post_by_file_FORM.py | 8 +- routes/legacy/post_by_url.py | 10 +- routes/v1/post_by_file_FORM.py | 4 +- routes/v1/post_by_url.py | 10 +- routes/ws_audio_transkrib.py | 10 +- tests/test_logging_config.py | 62 +++++++ user_instructions.md | 224 ++++++++++++++++++++++++ utils/do_logging.py | 33 +--- 14 files changed, 451 insertions(+), 62 deletions(-) create mode 100644 core/__init__.py create mode 100644 core/logging_config.py create mode 100644 tests/test_logging_config.py create mode 100644 user_instructions.md diff --git a/Recognizer/__init__.py b/Recognizer/__init__.py index 1263d0d..f559d56 100644 --- a/Recognizer/__init__.py +++ b/Recognizer/__init__.py @@ -1,3 +1,4 @@ +from starlette.requests import HTTPConnection import multiprocessing import numpy as np @@ -122,4 +123,6 @@ def apply_postprocessing(self, *params) -> list: logger.debug(f"Применяется постобработка текста для модели '{self.model_name}'") return self._post_processor(*params) -recognizer = Recognizer() + +def get_recognizer(conn: HTTPConnection): + return conn.app.state.recognizer diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index 4dc8848..ebeec24 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -1,5 +1,5 @@ +from fastapi import Depends import time - from pydub import AudioSegment from config import settings import asyncio @@ -14,7 +14,6 @@ from utils.do_logging import logger from utils.chunk_doing import find_last_speech_position from utils.resamppling import sync_resample_audiosegment -from Recognizer import recognizer from Recognizer.engine.stream_recognition import simple_recognise, recognise_w_speed_correction, simple_recognise_batch from Recognizer.engine.sentensizer import do_sensitizing from Recognizer.engine.echoe_clearing import remove_echo @@ -24,7 +23,7 @@ # Глобальный лок для потокобезопасности audio_lock = Lock() -def process_file(tmp_path, params): +def process_file(tmp_path, params, recognizer): process_file_start = time.perf_counter() res = False diarized = False @@ -122,7 +121,7 @@ def process_file(tmp_path, params): if params.use_batch: logger.info("Запрошен батчинг") - list_asr_result_wo_conf = simple_recognise_batch(audio_to_asr[post_id],params.batch_size) # --> list + list_asr_result_wo_conf = simple_recognise_batch(audio_to_asr[post_id],params.batch_size, recognizer) # --> list for _, asr_result_wo_conf in enumerate(list_asr_result_wo_conf): @@ -145,7 +144,7 @@ def process_file(tmp_path, params): params.speech_speed_correction_multiplier = multiplier else: # Производим распознавание - asr_result_wo_conf = asyncio.run(simple_recognise(audio_asr)) + asr_result_wo_conf = asyncio.run(simple_recognise(audio_asr, recognizer)) except Exception as e: logger.error(f"Error ASR audio - {e}") diff --git a/Recognizer/engine/stream_recognition.py b/Recognizer/engine/stream_recognition.py index b4ddc2d..223f668 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -1,7 +1,6 @@ import time import numpy as np from config import settings -from Recognizer import recognizer from utils.bytes_to_samples_audio import get_np_array_samples_float32 from utils.resamppling import sync_resample_audiosegment from utils.slow_down_audio import do_slow_down_audio @@ -41,7 +40,7 @@ def calc_speed(data): return speech_speed -async def simple_recognise(audio_data, ) -> dict: +async def simple_recognise(audio_data, recognizer) -> dict: # Приводим фреймрейт к фреймрейту модели if audio_data.frame_rate != settings.BASE_SAMPLE_RATE: audio_data = sync_resample_audiosegment(audio_data, settings.BASE_SAMPLE_RATE) @@ -86,7 +85,8 @@ async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_sl return result, speed, multiplier -def simple_recognise_batch(list_audio_data: list, batch_size: int = 8) -> list: +def simple_recognise_batch(list_audio_data: list, batch_size: int = 8, recognizer=None) -> list: + logger.info(f"Выполняется батчинг с размером {batch_size}") timer_sync_start = time.perf_counter() diff --git a/core/__init__.py b/core/__init__.py new file mode 100644 index 0000000..97daee7 --- /dev/null +++ b/core/__init__.py @@ -0,0 +1 @@ +# core package diff --git a/core/logging_config.py b/core/logging_config.py new file mode 100644 index 0000000..91b58b7 --- /dev/null +++ b/core/logging_config.py @@ -0,0 +1,97 @@ +import contextvars +import json +import logging +import sys +from datetime import datetime, timezone +from logging.config import dictConfig + +from config import settings + +# ContextVar для проброса request_id из middleware в лог-записи +request_id_var: contextvars.ContextVar[str | None] = contextvars.ContextVar("request_id", default=None) + + +class JsonFormatter(logging.Formatter): + """ + Форматтер, выводящий лог в виде JSON. + Поля: timestamp, level, logger, message, request_id, module, function, line. + """ + + def format(self, record: logging.LogRecord) -> str: + log_obj = { + "timestamp": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), + "level": record.levelname, + "logger": record.name, + "message": record.getMessage(), + "module": record.module, + "function": record.funcName, + "line": record.lineno, + "request_id": getattr(record, "request_id", None) or request_id_var.get(), + } + if record.exc_info: + log_obj["exception"] = self.formatException(record.exc_info) + return json.dumps(log_obj, ensure_ascii=False) + + +class RequestIDFilter(logging.Filter): + """Фильтр, пробрасывающий request_id из ContextVar в атрибут записи.""" + + def filter(self, record: logging.LogRecord) -> bool: + record.request_id = request_id_var.get() + return True + + +def setup_logging(): + """ + Централизованная настройка логирования. + """ + handlers = { + "stdout": { + "class": "logging.StreamHandler", + "stream": sys.stdout, + "formatter": "json", + "filters": ["request_id"], + }, + } + + if settings.IS_PROD: + handlers["file"] = { + "()": "logging.handlers.TimedRotatingFileHandler", + "filename": settings.FILENAME, + "when": "midnight", + "interval": 1, + "backupCount": settings.LOG_BACKUP_COUNT, + "encoding": "UTF-8", + "formatter": "json", + "filters": ["request_id"], + } + + handler_names = list(handlers.keys()) + + config = { + "version": 1, + "disable_existing_loggers": False, + "formatters": { + "json": { + "()": "core.logging_config.JsonFormatter", + }, + }, + "filters": { + "request_id": { + "()": "core.logging_config.RequestIDFilter", + }, + }, + "handlers": handlers, + "root": { + "level": settings.LOGGING_LEVEL, + "handlers": handler_names, + }, + "loggers": { + "uvicorn": {"level": "WARNING", "handlers": handler_names, "propagate": False}, + "uvicorn.access": {"level": "WARNING", "handlers": handler_names, "propagate": False}, + "httpx": {"level": "WARNING", "handlers": handler_names, "propagate": False}, + "httpcore": {"level": "WARNING", "handlers": handler_names, "propagate": False}, + }, + } + + dictConfig(config) diff --git a/main.py b/main.py index 562d66f..4f06187 100644 --- a/main.py +++ b/main.py @@ -1,4 +1,5 @@ -from utils.do_logging import logger +from Recognizer import Recognizer +import logging import uvicorn from config import settings import os @@ -15,6 +16,7 @@ from starlette.middleware.trustedhost import TrustedHostMiddleware from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware import uuid +from core.logging_config import setup_logging, request_id_var from utils.files_whatcher import start_file_watcher from utils.pre_start_init import paths import threading @@ -26,26 +28,38 @@ import models from config import WS_DESCRIPTION +logger = logging.getLogger(__name__) + + class RequestIDMiddleware(BaseHTTPMiddleware): async def dispatch(self, request: Request, call_next): request_id = request.headers.get("X-Request-ID", str(uuid.uuid4())) request.state.request_id = request_id - response = await call_next(request) - response.headers["X-Request-ID"] = request_id - return response + token = request_id_var.set(request_id) + try: + response = await call_next(request) + response.headers["X-Request-ID"] = request_id + return response + finally: + request_id_var.reset(token) @asynccontextmanager async def lifespan(app): + # Настройка логирования до любых других операций + setup_logging() + + # Установка HF_HOME для HuggingFace Hub + os.environ["HF_HOME"] = settings.HF_HOME + # on_start logger.debug("Приложение FastAPI запущено") - + # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) - - # Установка HF_HOME для HuggingFace Hub - os.environ["HF_HOME"] = settings.HF_HOME - + + app.state.recognizer = Recognizer() + if settings.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), @@ -55,6 +69,8 @@ async def lifespan(app): logger.info("File watcher started") yield # Здесь приложение работает + # cleanup (если нужно) + del app.state.recognize app = FastAPI( diff --git a/routes/legacy/post_by_file_FORM.py b/routes/legacy/post_by_file_FORM.py index 3d9143d..f53cdcd 100644 --- a/routes/legacy/post_by_file_FORM.py +++ b/routes/legacy/post_by_file_FORM.py @@ -5,16 +5,13 @@ from fastapi import APIRouter, Depends, File, Form, UploadFile from utils.do_logging import logger from models.fast_api_models import PostFileRequest, BaseResponse +from Recognizer import get_recognizer, Recognizer from Recognizer.engine.file_recognition import process_file from threading import Lock router = APIRouter() - -# Глобальный лок для потокобезопасности -audio_lock = Lock() - # Функция для извлечения параметров из FormData def get_file_request( keep_raw: bool = Form(default=True, description="Сохранять сырые данные."), @@ -48,6 +45,7 @@ def get_file_request( async def async_receive_file_legacy( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), + recognizer: Recognizer = Depends(get_recognizer) ) -> BaseResponse: # Сохраняем файл на диск асинхронно try: @@ -67,7 +65,7 @@ async def async_receive_file_legacy( logger.info(f"Получен и сохранён файл {file.filename}") try: # Запускаем обработку в потоке - result_dict = await asyncio.to_thread(process_file, buffer, params) + result_dict = await asyncio.to_thread(process_file, buffer, params, recognizer) result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/legacy/post_by_url.py b/routes/legacy/post_by_url.py index fdc2987..959467e 100644 --- a/routes/legacy/post_by_url.py +++ b/routes/legacy/post_by_url.py @@ -6,6 +6,10 @@ from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, BaseResponse + +from fastapi import Depends +from Recognizer import get_recognizer, Recognizer + from Recognizer.engine.file_recognition import process_file from threading import Lock from io import BytesIO @@ -17,7 +21,9 @@ audio_lock = Lock() @router.post("/post_one_step_req", response_model=BaseResponse) -async def post(params: SyncASRRequest) -> BaseResponse: +async def post(params: SyncASRRequest, + recognizer: Recognizer = Depends(get_recognizer) +) -> BaseResponse: """ На вход ждёт str(HttpUrl) - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n Если на вход передаётся не моно, то ответ будет в несколько элементов списка для каждого канала.\n @@ -54,7 +60,7 @@ async def post(params: SyncASRRequest) -> BaseResponse: try: # Запускаем обработку в потоке - result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) + result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params, recognizer) result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 9e6dc0e..3d980b0 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -5,6 +5,7 @@ from fastapi import APIRouter, Depends, File, Form, UploadFile from models.fast_api_models import PostFileRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file +from Recognizer import get_recognizer, Recognizer router = APIRouter() # Функция для извлечения параметров из FormData @@ -40,6 +41,7 @@ def get_file_request( async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), + recognizer: Recognizer = Depends(get_recognizer) ) -> V1ASRResponse: try: buffer = BytesIO(await file.read()) @@ -55,7 +57,7 @@ async def async_receive_file( else: logger.info(f"Получен и сохранён файл {file.filename}") try: - result_dict = await asyncio.to_thread(process_file, buffer, params) + result_dict = await asyncio.to_thread(process_file, buffer, params, recognizer) return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/routes/v1/post_by_url.py b/routes/v1/post_by_url.py index 6293e6c..b55c3d4 100644 --- a/routes/v1/post_by_url.py +++ b/routes/v1/post_by_url.py @@ -5,12 +5,18 @@ from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData + +from fastapi import Depends +from Recognizer import get_recognizer, Recognizer + from Recognizer.engine.file_recognition import process_file + router = APIRouter() @router.post("/post_one_step_req", response_model=V1ASRResponse) -async def post_v1(params: SyncASRRequest) -> V1ASRResponse: +async def post_v1(params: SyncASRRequest, + recognizer: Recognizer = Depends(get_recognizer)) -> V1ASRResponse: post_id = uuid.uuid4() if params.AudioFileUrl: res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) @@ -26,7 +32,7 @@ async def post_v1(params: SyncASRRequest) -> V1ASRResponse: ) try: - result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) + result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params, recognizer) return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 0458a32..685f224 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -13,6 +13,9 @@ from utils.tokens_to_Result import process_single_token_vocab_output from utils.resamppling import async_resample_audiosegment +from fastapi import Depends +from Recognizer import get_recognizer, Recognizer + from Recognizer.engine.sentensizer import do_sensitizing from Recognizer.engine.stream_recognition import simple_recognise @@ -20,7 +23,8 @@ router = APIRouter() @router.websocket("/ws") -async def websocket(ws: WebSocket): +async def websocket(ws: WebSocket, + recognizer: Recognizer = Depends(get_recognizer)): wait_null_answers=True client_id = uuid.uuid4() logger.debug(f'Принят новый сокет id = {client_id}') @@ -116,7 +120,7 @@ async def websocket(ws: WebSocket): logger.error(f"AcceptWaveform error - {e} in channel {channel_name}") else: try: - asr_result = await simple_recognise(audio_to_asr[client_id][-1]) + asr_result = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) asr_result_words = process_single_token_vocab_output(asr_result, audio_duration[client_id]) audio_duration[client_id] += audio_to_asr[client_id][-1].duration_seconds logger.debug(asr_result_words) @@ -190,7 +194,7 @@ async def websocket(ws: WebSocket): last_result = None error_description = f"Ошибка дополнения тишиной последнего чанка - {e} in channel {channel_name}" else: - last_asr_result_w_conf = await simple_recognise(audio_to_asr[client_id][-1]) + last_asr_result_w_conf = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) last_result = process_single_token_vocab_output(last_asr_result_w_conf, audio_duration[client_id]) logger.debug(f'Последний результат {last_result.get("data").get("text")} in channel {channel_name}') diff --git a/tests/test_logging_config.py b/tests/test_logging_config.py new file mode 100644 index 0000000..33ab3ef --- /dev/null +++ b/tests/test_logging_config.py @@ -0,0 +1,62 @@ +import json +import logging +import pytest +from core.logging_config import JsonFormatter, setup_logging, request_id_var + + +def test_json_formatter_outputs_valid_json(): + """JsonFormatter должен возвращать валидную JSON-строку с обязательными полями.""" + formatter = JsonFormatter() + record = logging.LogRecord( + name="test", level=logging.INFO, pathname="test.py", lineno=1, + msg="hello world", args=(), exc_info=None, func="" + ) + output = formatter.format(record) + parsed = json.loads(output) + + assert parsed["message"] == "hello world" + assert parsed["level"] == "INFO" + assert parsed["logger"] == "test" + assert parsed["module"] == "test" + assert parsed["function"] == "" + assert parsed["line"] == 1 + assert "timestamp" in parsed + assert "request_id" in parsed + + +def test_json_formatter_includes_request_id_from_contextvar(): + """При установленном ContextVar request_id должен попадать в лог.""" + formatter = JsonFormatter() + token = request_id_var.set("req-abc-123") + try: + record = logging.LogRecord( + name="test", level=logging.DEBUG, pathname="test.py", lineno=2, + msg="with request id", args=(), exc_info=None, func="test_func" + ) + output = formatter.format(record) + parsed = json.loads(output) + assert parsed["request_id"] == "req-abc-123" + finally: + request_id_var.reset(token) + + +def test_setup_logging_configures_root_logger(): + """setup_logging() должен создать хотя бы один handler у root-логгера.""" + for handler in logging.root.handlers[:]: + logging.root.removeHandler(handler) + + setup_logging() + + assert len(logging.root.handlers) > 0 + handler = logging.root.handlers[0] + assert isinstance(handler.formatter, JsonFormatter) + + +def test_do_logging_import_has_no_side_effects(): + """Импорт utils.do_logging не должен модифицировать root handlers.""" + before = logging.root.handlers.copy() + import importlib + import utils.do_logging + importlib.reload(utils.do_logging) + after = logging.root.handlers.copy() + assert before == after diff --git a/user_instructions.md b/user_instructions.md new file mode 100644 index 0000000..aab8137 --- /dev/null +++ b/user_instructions.md @@ -0,0 +1,224 @@ +# Инструкции по настройке и эксплуатации ASR FastAPI + +## 1. Общие принципы + +Все параметры приложения задаются через переменные окружения или файл `.env` в корне проекта. +Приоритет: переменные окружения > `.env` > значения по умолчанию в `config.py`. + +## 2. RequestIDMiddleware (X-Request-ID) + +### Назначение +Каждый HTTP-запрос и ответ снабжается уникальным идентификатором `X-Request-ID`. Это позволяет: +- связывать логи на балансировщике, сервере приложений и клиенте; +- отслеживать цепочку вызовов при отладке; +- идентифицировать повторные запросы (retry) от клиента. + +### Как работает +1. Клиент **может** отправить собственный `X-Request-ID` в заголовке запроса. +2. Если заголовок отсутствует, middleware автоматически генерирует UUID4. +3. Идентификатор сохраняется в `request.state.request_id` и доступен внутри любого роута: +``` + request.state.request_id +``` +4. В любом случае заголовок `X-Request-ID` возвращается в ответе клиенту. + +### Пример использования клиентом +```bash +# Передаём свой идентификатор +curl -H "X-Request-ID: trace-abc-123" http://localhost:49153/ + +# Ответ содержит тот же X-Request-ID +``` + +### Рекомендации +- В микросервисной архитектуре пробрасывайте `X-Request-ID` дальше во все downstream-запросы. +- Используйте `request.state.request_id` при записи логов внутри роутов и сервисов. +- Не генерируйте ID на клиенте, если не требуется сквозная трассировка — сервер создаст его сам. + +## 3. TrustedHostMiddleware + +### Назначение +Отклоняет HTTP-запросы с заголовком `Host`, не входящим в белый список. Защищает от DNS Rebinding и фальшивых заголовков. + +### Параметр конфигурации +- `ALLOWED_HOSTS` — список разрешённых хостов. + +### Примеры значений + +**Разработка:** +```bash +ALLOWED_HOSTS=["*"] +``` + +**Продакшен (за reverse-proxy):** +```bash +ALLOWED_HOSTS=["asr.example.com","api.example.com"] +``` + +### Проверка +```bash +curl -H "Host: evil.com" http://localhost:49153/ +# Ожидается: 400 Bad Request +``` + +## 4. GZipMiddleware + +### Назначение +Сжимает HTTP-ответы размером более 1000 байт для экономии трафика. + +### Настройка +Параметр `minimum_size=1000` задан в коде и не требует конфигурации через `.env`. +При необходимости измените значение в `main.py`. + +### Проверка +```bash +curl -H "Accept-Encoding: gzip" http://localhost:49153/openapi.json +# В ответе должен присутствовать заголовок content-encoding: gzip +``` + +## 5. CORS (Cross-Origin Resource Sharing) + +### Назначение +Разрешает браузерным клиентам (фронтенду) обращаться к API с другого домена. + +### Параметр конфигурации +- `CORS_ORIGINS` — список разрешённых доменов (JSON-список строк). + +### Примеры значений + +**Разработка:** +```bash +CORS_ORIGINS=["http://localhost:3000","http://127.0.0.1:8080"] +``` + +**Продакшен:** +```bash +CORS_ORIGINS=["https://asr.example.com","https://admin.example.com"] +``` + +### ⚠️ Важное предупреждение по безопасности +- **Никогда** не используйте `CORS_ORIGINS=["*"]` в продакшене совместно с `allow_credentials=True`. +- При `IS_PROD=True` и `CORS_ORIGINS=["*"]` приложение выдаст `WARNING` в логи. +- Если фронтенд работает на другом домене, всегда указывайте его явно. + +## 6. Комбинированная настройка безопасности для продакшена + +Рекомендуемый минимальный набор параметров `.env` для продакшена: + +```bash +IS_PROD=true +ALLOWED_HOSTS=["asr.example.com"] +CORS_ORIGINS=["https://admin.example.com","https://app.example.com"] +``` + +## 7. Запуск и тестирование + +### Запуск сервера +```bash +python main.py +# или +uvicorn main:app --host 0.0.0.0 --port 49153 +``` + +### Запуск тестов (без поднятия сервера) +```bash +python -m pytest tests/ -v +``` + +## 8. Подготовка фронтенда + +При разработке фронтенда убедитесь, что: +1. Домен фронтенда добавлен в `CORS_ORIGINS`. +2. Фронтенд отправляет запросы с правильным заголовком `Host` (если используется TrustedHost). +3. При использовании WebSocket (`/ws`) учитывайте, что CORS-политика не распространяется на WebSocket напрямую; проверяйте `Origin` на стороне сервера при необходимости. +4. Для трассировки запросов передавайте и сохраняйте заголовок `X-Request-ID`. + +--- + +## 9. ProxyHeadersMiddleware + +### Назначение +Когда приложение работает за reverse-proxy (Nginx, Traefik, Kubernetes Ingress, CloudFlare), все запросы приходят с IP-адреса прокси. `ProxyHeadersMiddleware` позволяет доверять заголовкам `X-Forwarded-For`, +`X-Forwarded-Proto`, `X-Forwarded-Port` и подменять `request.client.host` / `request.url.scheme` на реальные значения клиента. + +### Параметр конфигурации +- `TRUSTED_PROXIES` — список IP-адресов или CIDR прокси-серверов, которым можно доверять. `["*"]` означает доверие любому (только если прокси гарантирует очистку заголовков!). + +### Примеры значений + +**Разработка (без прокси):** +```bash +TRUSTED_PROXIES=["*"] +``` + +**Продакшен (конкретные прокси):** +```bash +TRUSTED_PROXIES=["10.0.0.0/8","172.16.0.0/12","127.0.0.1"] +``` + +### ⚠️ Важное предупреждение по безопасности +- **Никогда** не используйте `TRUSTED_PROXIES=["*"]`, если ваш сервер доступен напрямую из интернета без прокси. Злоумышленник сможет подделать `X-Forwarded-For`. +- Настраивайте список только с IP-адресов вашего балансировщика/прокси. +- Убедитесь, что прокси очищает входящие `X-Forwarded-*` заголовки от клиентов (например, `real_ip_header` в Nginx). + +### Как работает +1. Запрос приходит от прокси с заголовком `X-Forwarded-For: <реальный_IP_клиента>`. +2. Middleware проверяет, что IP прокси (из `scope["client"]`) входит в `TRUSTED_PROXIES`. +3. Если да, `request.client.host` заменяется на IP из `X-Forwarded-For`. +4. Аналогично `X-Forwarded-Proto` обновляет `request.url.scheme` (http → https). + +### Проверка +```bash +# Если TRUSTED_PROXIES=["*"], запрос с X-Forwarded-For должен пройти +curl -H "X-Forwarded-For: 203.0.113.1" http://localhost:49153/ +``` + +--- + +======= +## 10. Логирование (JSON) + +### Назначение +Все логи приложения выводятся в структурированном JSON-формате. Это позволяет собирать, фильтровать и анализировать логи в централизованных системах (Kibana, Grafana Loki, Fluent Bit). + +### Как работает +1. Конфигурация находится в `core/logging_config.py`. +2. Функция `setup_logging()` вызывается в `lifespan` при старте приложения. +3. `JsonFormatter` формирует записи с полями: `timestamp`, `level`, `logger`, `message`, `request_id`, `module`, `function`, `line`. +4. `RequestIDMiddleware` автоматически пробрасывает `X-Request-ID` в лог-записи через `ContextVar`. + +### Параметр конфигурации +- `LOGGING_LEVEL` — уровень логирования (`DEBUG`, `INFO`, `WARNING`, `ERROR`). +- `LOG_BACKUP_COUNT` — количество дней хранения логов (только в `IS_PROD=true`). +- `FILENAME` — путь к файлу логов (только в `IS_PROD=true`). + +### Примеры значений + +**Разработка:** +```bash +IS_PROD=false +LOGGING_LEVEL=DEBUG +``` + +**Продакшен:** +```bash +IS_PROD=true +LOGGING_LEVEL=INFO +LOG_BACKUP_COUNT=30 +``` + +### Проверка +```bash +# После запуска сервера в stdout должны идти JSON-строки +curl http://localhost:49153/ +# {"timestamp": "2025-01-15T12:34:56.789Z", "level": "INFO", ...} +``` + +### Рекомендации +- В продакшене логи пишутся одновременно в `stdout` (для Docker/k8s) и в ротируемый файл. +- Уровень `uvicorn.access` и `httpx` понижен до `WARNING`, чтобы не засорять логи. +- Для трассировки используйте поле `request_id` — оно связывает все логи одного HTTP-запроса. + +--- + +*Документ обновляется по мере добавления нового функционала (middleware, auth, metrics).* \ No newline at end of file diff --git a/utils/do_logging.py b/utils/do_logging.py index ff418a5..76e1463 100644 --- a/utils/do_logging.py +++ b/utils/do_logging.py @@ -1,35 +1,6 @@ # -*- coding: utf-8 -*- -from config import settings - import logging -from logging.handlers import TimedRotatingFileHandler -from fastapi.logger import logger as fastapi_logger +# Логгер модуля. Полная конфигурация производится в lifespan через core.logging_config.setup_logging() +# Этот модуль не должен содержать side-effects при импорте. logger = logging.getLogger(__name__) - -if settings.IS_PROD: - # Создаем обработчик, который ротирует логи каждый день в полночь - file_handler = TimedRotatingFileHandler( - filename=settings.FILENAME, # Базовое имя файла - when='midnight', # Ротация каждый день в полночь - interval=1, # Интервал - каждый день - backupCount=settings.LOG_BACKUP_COUNT if hasattr(settings, 'LOG_BACKUP_COUNT') else 7, - # Хранить 7 дней логов по умолчанию - encoding='UTF-8' - ) - - file_handler.setLevel(settings.LOGGING_LEVEL) - file_handler.setFormatter(logging.Formatter(settings.LOGGING_FORMAT)) - - # Настраиваем логгер - logger.addHandler(file_handler) - fastapi_logger.addHandler(file_handler) - - # Убираем basicConfig, так как мы используем кастомный обработчик - logging.basicConfig(level=settings.LOGGING_LEVEL) -else: - logging.basicConfig( - level=settings.LOGGING_LEVEL, - format=settings.LOGGING_FORMAT, - encoding="UTF-8" - ) \ No newline at end of file From a8eee8898888b487542c2bb57f038f0353e4ab86 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 30 Apr 2026 15:19:02 +0300 Subject: [PATCH 16/46] =?UTF-8?q?=D0=B8=D1=81=D0=BF=D0=BE=D0=BB=D1=8C?= =?UTF-8?q?=D0=B7=D0=BE=D0=B2=D0=B0=D0=BD=D0=B8=D0=B5=20=D0=B3=D0=BB=D0=BE?= =?UTF-8?q?=D0=B1=D0=B0=D0=BB=D1=8C=D0=BD=D0=BE=D0=B3=D0=BE=20=D0=BB=D0=BE?= =?UTF-8?q?=D0=B3=D0=B3=D0=B5=D1=80=D0=B0.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Diarisation/__init__.py | 4 +- Diarisation/diarazer.py | 5 +- Diarisation/do_diarize.py | 5 +- Punctuation/__init__.py | 249 ++++++++++++++++++++++-- Punctuation/punctuate.py | 233 ---------------------- Recognizer/__init__.py | 3 +- Recognizer/engine/echoe_clearing.py | 5 +- Recognizer/engine/file_recognition.py | 15 +- Recognizer/engine/sentensizer.py | 14 +- Recognizer/engine/stream_recognition.py | 3 +- VoiceActivityDetector/__init__.py | 4 +- VoiceActivityDetector/do_vad.py | 5 +- main.py | 14 +- routes/legacy/demo_page.py | 4 +- routes/legacy/post_by_file_FORM.py | 17 +- routes/legacy/post_by_url.py | 12 +- routes/v1/post_by_file_FORM.py | 15 +- routes/v1/post_by_url.py | 11 +- routes/ws_audio_transkrib.py | 12 +- utils/chunk_doing.py | 3 +- utils/file_exists.py | 3 +- utils/save_asr_to_md.py | 10 +- utils/send_messages.py | 4 +- utils/slow_down_audio.py | 4 +- utils/tokens_to_Result.py | 8 +- 25 files changed, 352 insertions(+), 310 deletions(-) delete mode 100644 Punctuation/punctuate.py diff --git a/Diarisation/__init__.py b/Diarisation/__init__.py index e1c4841..79f12dc 100644 --- a/Diarisation/__init__.py +++ b/Diarisation/__init__.py @@ -1,9 +1,9 @@ from config import settings from utils.pre_start_init import paths -from utils.do_logging import logger from VoiceActivityDetector import vad import requests - +import logging +logger = logging.getLogger(__name__) if settings.CAN_DIAR: if not paths.get("diar_speaker_model_path").exists(): diff --git a/Diarisation/diarazer.py b/Diarisation/diarazer.py index fb1c439..091446b 100644 --- a/Diarisation/diarazer.py +++ b/Diarisation/diarazer.py @@ -3,13 +3,14 @@ from config import settings from Diarisation.do_diarize import load_and_preprocess_audio from utils.pre_start_init import posted_and_downloaded_audio -from utils.do_logging import logger - from collections import defaultdict +import logging +logger = logging.getLogger(__name__) if settings.CAN_DIAR: from Diarisation import diarizer + async def do_diarizing( file_id:str, asr_raw_data, diff --git a/Diarisation/do_diarize.py b/Diarisation/do_diarize.py index a642d61..ec7b821 100644 --- a/Diarisation/do_diarize.py +++ b/Diarisation/do_diarize.py @@ -13,10 +13,9 @@ import time from umap import UMAP from hdbscan import HDBSCAN -from utils.do_logging import logger -from utils.pre_start_init import paths from utils.resamppling import async_resample_audiosegment - +import logging +logger = logging.getLogger(__name__) class Diarizer: def __init__(self, embedding_model_path: str, diff --git a/Punctuation/__init__.py b/Punctuation/__init__.py index c70548b..bb086be 100644 --- a/Punctuation/__init__.py +++ b/Punctuation/__init__.py @@ -1,11 +1,238 @@ -from utils.pre_start_init import paths -from utils.do_logging import logger -from .punctuate import SbertPuncCaseOnnx -from config import settings - -try: - sbertpunc = SbertPuncCaseOnnx(paths.get("punctuation_model_path"), use_gpu = settings.PUNCTUATE_WITH_GPU) -except Exception as e: - logger.error(f"Error getting punctuation model - {e}") -else: - logger.info(f'Успешно загружена модель Пунктуации') \ No newline at end of file +# -*- coding: utf-8 -*- +from starlette.requests import HTTPConnection +import logging +import asyncio +import numpy as np +from transformers import AutoTokenizer +import onnxruntime as ort +from pathlib import Path +import pynvml + +from typing import List +logger = logging.getLogger(__name__) + +# Прогнозируемые знаки препинания +PUNK_MAPPING = {".": "PERIOD", ",": "COMMA", "?": "QUESTION"} + +# Прогнозируемый регистр LOWER - нижний регистр, UPPER - верхний регистр для первого символа, +# UPPER_TOTAL - верхний регистр для всех символов +LABELS_CASE = ["LOWER", "UPPER", "UPPER_TOTAL"] +# Добавим в пунктуацию метку O означающий отсутствие пунктуации +LABELS_PUNC = ["O"] + list(PUNK_MAPPING.values()) + +# Сформируем метки на основе комбинаций регистра и пунктуации +LABELS_list = [] + +for case in LABELS_CASE: + for punc in LABELS_PUNC: + LABELS_list.append(f"{case}_{punc}") + +LABELS = {label: i + 1 for i, label in enumerate(LABELS_list)} +LABELS["O"] = -100 +INVERSE_LABELS = {i: label for label, i in LABELS.items()} + +LABEL_TO_PUNC_LABEL = { + label: label.split("_")[-1] for label in LABELS.keys() if label != "O" +} +LABEL_TO_CASE_LABEL = { + label: "_".join(label.split("_")[:-1]) for label in LABELS.keys() if label != "O" +} + + +def token_to_label(token, label): + if type(label) == int: + label = INVERSE_LABELS[label] + if label == "LOWER_O": + return token + if label == "LOWER_PERIOD": + return token + "." + if label == "LOWER_COMMA": + return token + "," + if label == "LOWER_QUESTION": + return token + "?" + if label == "UPPER_O": + return token.capitalize() + if label == "UPPER_PERIOD": + return token.capitalize() + "." + if label == "UPPER_COMMA": + return token.capitalize() + "," + if label == "UPPER_QUESTION": + return token.capitalize() + "?" + if label == "UPPER_TOTAL_O": + return token.upper() + if label == "UPPER_TOTAL_PERIOD": + return token.upper() + "." + if label == "UPPER_TOTAL_COMMA": + return token.upper() + "," + if label == "UPPER_TOTAL_QUESTION": + return token.upper() + "?" + if label == "O": + return token + + +def decode_label(label, classes="all"): + if classes == "punc": + return LABEL_TO_PUNC_LABEL[INVERSE_LABELS[label]] + if classes == "case": + return LABEL_TO_CASE_LABEL[INVERSE_LABELS[label]] + else: + return INVERSE_LABELS[label] + + +class SbertPuncCaseOnnx: + def __init__(self, onnx_model_path, use_gpu = False): + self.sessions: List[ort.InferenceSession] = [] + + self.tokenizer = AutoTokenizer.from_pretrained(onnx_model_path, + strip_accents=False, + ) + session_options = ort.SessionOptions() + session_options.log_severity_level = 4 # Выключаем подробный лог + session_options.enable_profiling = False + session_options.enable_mem_pattern = False # True в диаризации + session_options.enable_mem_reuse = False # True в диаризации + session_options.enable_cpu_mem_arena = False # True в диаризации + session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL + session_options.inter_op_num_threads = 0 + session_options.intra_op_num_threads = 0 + session_options.add_session_config_entry("session.disable_prepacking", "1") # Отключаем дублирование весов + session_options.add_session_config_entry("session.use_device_allocator_for_initializers", "1") + + if not use_gpu: + providers = ['CPUExecutionProvider'] + else: + providers = [('CUDAExecutionProvider', { + 'device_id': 0, + 'arena_extend_strategy': 'kSameAsRequested', # 20.657809 на 1000 итераций и 26 Мб съел + 'gpu_mem_limit': int(1.8 * 1024 * 1024 * 1024), # Потребляет где-то 1.9 Гб памяти ГПУ + 'cudnn_conv_algo_search': 'EXHAUSTIVE', + 'do_copy_in_default_stream': True, + }), + 'CPUExecutionProvider'] + + model_pth = Path(onnx_model_path) / "model.onnx" + + # Todo - можно собрать очередь сессий. Очень интересный механизм для оптимизации производительности + self.session = ort.InferenceSession(path_or_bytes=model_pth, + sess_options=session_options, + providers=providers, + logger=logger) + + + async def punctuate(self, text): + text = text.strip().lower() + # Разобьем предложение на слова + words = text.split() + + tokenizer_output = self.tokenizer(words, is_split_into_words=True) + + if len(tokenizer_output.input_ids) > 512: + return " ".join( + [ + await self.punctuate(" ".join(text_part)) + for text_part in np.array_split(words, 2) + ] + ) + + # Подготовка входных данных для модели + input_ids = np.array(tokenizer_output.input_ids, dtype=np.int64).reshape(1, -1) + attention_mask = np.array(tokenizer_output.attention_mask, dtype=np.int64).reshape(1, -1) + token_type_ids = np.zeros_like(input_ids, dtype=np.int64) # Добавляем token_type_ids + + # Выполнение модели + loop = asyncio.get_event_loop() + outputs = await loop.run_in_executor( + None, + lambda: self.session.run(None, { + "input_ids": input_ids, + "attention_mask": attention_mask, + "token_type_ids": token_type_ids, # Передаём token_type_ids + } )) + + predictions = np.argmax(outputs[0], axis=2) + + # decode punctuation and casing + splitted_text = [] + word_ids = tokenizer_output.word_ids() + for i, word in enumerate(words): + label_pos = word_ids.index(i) + label_id = predictions[0][label_pos] + label = decode_label(label_id) + splitted_text.append(token_to_label(word, label)) + capitalized_text = " ".join(splitted_text) + return capitalized_text + + async def process_punctuation_sessions(self, text): + try: + capitalized_text = await self.punctuate(text) + except Exception as e: + logger.error(f"Ошибка пунктуатора - {e}") + + return capitalized_text + +def gpu_stat(gpu_index): + + try: + pynvml.nvmlInit() + handle = pynvml.nvmlDeviceGetHandleByIndex(gpu_index) # Первая видеокарта + mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle) + free_mb = mem_info.free / 1024**2 + utilization = pynvml.nvmlDeviceGetUtilizationRates(handle) + gpu_load = utilization.gpu + temperature = pynvml.nvmlDeviceGetTemperature(handle, pynvml.NVML_TEMPERATURE_GPU) + except pynvml.NVMLError as e: + return {"error": str(e)} + finally: + pynvml.nvmlShutdown() + return free_mb, gpu_load,temperature + + +def get_punctuator(conn: HTTPConnection): + return conn.app.state.punctuator + + +if __name__ == '__main__': + from datetime import datetime as dt + + test_input_text = [ + "Живут грызуны по лесам и полям что они едят эти зверьки грызут зерна и кору деревьев зайцы обгладывают яблони в саду грызуны портят посевы", + "Солнце только что встало небо ясное все вокруг блестит как хорошо на свежем воздухе слышишь пение жаворонка звонкий голосок слышен в ясной вышине", + "В сырых местах живут неядовитые змейки у них два желтых пятна на затылке ужи любят воду и хорошо плавают кормятся ужи лягушками и рыбами вы видели ужа не бойтесь его", + "Белка бойко лазит по деревьям какая она ловкая на ушах у белки кисточки хвост длинный и пушистый зачем он ей белка прикрывается хвостом от холода он служит ей рулём при прыжках", + "Когда ты был в цирке вспомни яркие афиши и флажки вот жонглёр ловит на лету тарелки фокусник превратил шляпу в букет цветов а вот клоун он смешит людей у него в куртке петух как интересно в цирке", + "Вот ночная хищная птица голова у нее круглая клюв крючком когти острые Узнали ее это сова она живет в лесах или на чердаках домов ночью птица ловит мышей", + "Черепахи живут на земле и в воде они откладывают яйца прямо на камни черепахи не высиживают их яйца лопаются сами появляются маленькие черепашата а какого размера эти пресмыкающиеся черепахи бывают маленькие и очень большие", + ] + + model_path = str(Path("../models/sbert_punc_case_ru_onnx")) + print(f" ресурсы до старта приложения {gpu_stat(0)}") + time_start = dt.now() + sbertpunc = SbertPuncCaseOnnx(model_path, use_gpu=True) + print(f"Время на инициализацию {(dt.now() - time_start).total_seconds()}") + print(f"Ресурсы после старта приложения {gpu_stat(0)}") + import logging as logger + + # for _ , text_ in enumerate(input_text*100): + # punctuated = asyncio.run(sbertpunc.process(text_) + # ) + # # print(punctuated) + + async def process_texts(texts): + tasks = [sbertpunc.process_punctuation_sessions(text) for text in texts] + results = await asyncio.gather(*tasks, return_exceptions=True) + return results + + + time_start = dt.now() + texts = test_input_text * 100 + punctuated_texts = asyncio.run(process_texts(texts)) + # for text, punctuated in zip(texts, punctuated_texts): + # print(f"Punctuated: {punctuated}\n") + + + + print(f" ресурсы после окончания работы приложения {gpu_stat(0)}") + print(f"Время выполнения {(dt.now() - time_start).total_seconds()}") + +### Увеличение количества сессий в адаптере приводит к пропорциональному увеличению расхода памяти (1,9 * Х для пунктуации) и где-то на 30% быстрее 7,676 против 5,579. +### Если использовать 2 ГПУ, то выполняется где-то на быстрее 60% 4.63946 против 7,676 \ No newline at end of file diff --git a/Punctuation/punctuate.py b/Punctuation/punctuate.py deleted file mode 100644 index 0284531..0000000 --- a/Punctuation/punctuate.py +++ /dev/null @@ -1,233 +0,0 @@ -# -*- coding: utf-8 -*- -import asyncio -import datetime -import logging as logger - -import numpy as np -from transformers import AutoTokenizer -import onnxruntime as ort -from pathlib import Path -import pynvml - -from typing import List - -# Прогнозируемые знаки препинания -PUNK_MAPPING = {".": "PERIOD", ",": "COMMA", "?": "QUESTION"} - -# Прогнозируемый регистр LOWER - нижний регистр, UPPER - верхний регистр для первого символа, -# UPPER_TOTAL - верхний регистр для всех символов -LABELS_CASE = ["LOWER", "UPPER", "UPPER_TOTAL"] -# Добавим в пунктуацию метку O означающий отсутствие пунктуации -LABELS_PUNC = ["O"] + list(PUNK_MAPPING.values()) - -# Сформируем метки на основе комбинаций регистра и пунктуации -LABELS_list = [] - -for case in LABELS_CASE: - for punc in LABELS_PUNC: - LABELS_list.append(f"{case}_{punc}") - -LABELS = {label: i + 1 for i, label in enumerate(LABELS_list)} -LABELS["O"] = -100 -INVERSE_LABELS = {i: label for label, i in LABELS.items()} - -LABEL_TO_PUNC_LABEL = { - label: label.split("_")[-1] for label in LABELS.keys() if label != "O" -} -LABEL_TO_CASE_LABEL = { - label: "_".join(label.split("_")[:-1]) for label in LABELS.keys() if label != "O" -} - - -def token_to_label(token, label): - if type(label) == int: - label = INVERSE_LABELS[label] - if label == "LOWER_O": - return token - if label == "LOWER_PERIOD": - return token + "." - if label == "LOWER_COMMA": - return token + "," - if label == "LOWER_QUESTION": - return token + "?" - if label == "UPPER_O": - return token.capitalize() - if label == "UPPER_PERIOD": - return token.capitalize() + "." - if label == "UPPER_COMMA": - return token.capitalize() + "," - if label == "UPPER_QUESTION": - return token.capitalize() + "?" - if label == "UPPER_TOTAL_O": - return token.upper() - if label == "UPPER_TOTAL_PERIOD": - return token.upper() + "." - if label == "UPPER_TOTAL_COMMA": - return token.upper() + "," - if label == "UPPER_TOTAL_QUESTION": - return token.upper() + "?" - if label == "O": - return token - - -def decode_label(label, classes="all"): - if classes == "punc": - return LABEL_TO_PUNC_LABEL[INVERSE_LABELS[label]] - if classes == "case": - return LABEL_TO_CASE_LABEL[INVERSE_LABELS[label]] - else: - return INVERSE_LABELS[label] - - -class SbertPuncCaseOnnx: - def __init__(self, onnx_model_path, use_gpu = False): - self.sessions: List[ort.InferenceSession] = [] - - self.tokenizer = AutoTokenizer.from_pretrained(onnx_model_path, - strip_accents=False, - ) - session_options = ort.SessionOptions() - session_options.log_severity_level = 4 # Выключаем подробный лог - session_options.enable_profiling = False - session_options.enable_mem_pattern = False # True в диаризации - session_options.enable_mem_reuse = False # True в диаризации - session_options.enable_cpu_mem_arena = False # True в диаризации - session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL - session_options.inter_op_num_threads = 0 - session_options.intra_op_num_threads = 0 - session_options.add_session_config_entry("session.disable_prepacking", "1") # Отключаем дублирование весов - session_options.add_session_config_entry("session.use_device_allocator_for_initializers", "1") - - if not use_gpu: - providers = ['CPUExecutionProvider'] - else: - providers = [('CUDAExecutionProvider', { - 'device_id': 0, - 'arena_extend_strategy': 'kSameAsRequested', # 20.657809 на 1000 итераций и 26 Мб съел - 'gpu_mem_limit': int(1.8 * 1024 * 1024 * 1024), # Потребляет где-то 1.9 Гб памяти ГПУ - 'cudnn_conv_algo_search': 'EXHAUSTIVE', - 'do_copy_in_default_stream': True, - }), - 'CPUExecutionProvider'] - - model_pth = Path(onnx_model_path) / "model.onnx" - - # Todo - можно собрать очередь сессий. Очень интересный механизм для оптимизации производительности - self.session = ort.InferenceSession(path_or_bytes=model_pth, - sess_options=session_options, - providers=providers) - - - async def punctuate(self, text): - text = text.strip().lower() - # Разобъем предложение на слова - words = text.split() - - tokenizer_output = self.tokenizer(words, is_split_into_words=True) - - if len(tokenizer_output.input_ids) > 512: - return " ".join( - [ - await self.punctuate(" ".join(text_part)) - for text_part in np.array_split(words, 2) - ] - ) - - # Подготовка входных данных для модели - input_ids = np.array(tokenizer_output.input_ids, dtype=np.int64).reshape(1, -1) - attention_mask = np.array(tokenizer_output.attention_mask, dtype=np.int64).reshape(1, -1) - token_type_ids = np.zeros_like(input_ids, dtype=np.int64) # Добавляем token_type_ids - - # Выполнение модели - loop = asyncio.get_event_loop() - outputs = await loop.run_in_executor( - None, - lambda: self.session.run(None, { - "input_ids": input_ids, - "attention_mask": attention_mask, - "token_type_ids": token_type_ids, # Передаём token_type_ids - } )) - - predictions = np.argmax(outputs[0], axis=2) - - # decode punctuation and casing - splitted_text = [] - word_ids = tokenizer_output.word_ids() - for i, word in enumerate(words): - label_pos = word_ids.index(i) - label_id = predictions[0][label_pos] - label = decode_label(label_id) - splitted_text.append(token_to_label(word, label)) - capitalized_text = " ".join(splitted_text) - return capitalized_text - - async def process_punctuation_sessions(self, text): - try: - capitalized_text = await self.punctuate(text) - except Exception as e: - logger.error(f"Ошибка пунктуатора - {e}") - - return capitalized_text - -def gpu_stat(gpu_index): - - try: - pynvml.nvmlInit() - handle = pynvml.nvmlDeviceGetHandleByIndex(gpu_index) # Первая видеокарта - mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle) - free_mb = mem_info.free / 1024**2 - utilization = pynvml.nvmlDeviceGetUtilizationRates(handle) - gpu_load = utilization.gpu - temperature = pynvml.nvmlDeviceGetTemperature(handle, pynvml.NVML_TEMPERATURE_GPU) - except pynvml.NVMLError as e: - return {"error": str(e)} - finally: - pynvml.nvmlShutdown() - return free_mb, gpu_load,temperature - - -if __name__ == '__main__': - from datetime import datetime as dt - - test_input_text = [ - "Живут грызуны по лесам и полям что они едят эти зверьки грызут зерна и кору деревьев зайцы обгладывают яблони в саду грызуны портят посевы", - "Солнце только что встало небо ясное все вокруг блестит как хорошо на свежем воздухе слышишь пение жаворонка звонкий голосок слышен в ясной вышине", - "В сырых местах живут неядовитые змейки у них два желтых пятна на затылке ужи любят воду и хорошо плавают кормятся ужи лягушками и рыбами вы видели ужа не бойтесь его", - "Белка бойко лазит по деревьям какая она ловкая на ушах у белки кисточки хвост длинный и пушистый зачем он ей белка прикрывается хвостом от холода он служит ей рулём при прыжках", - "Когда ты был в цирке вспомни яркие афиши и флажки вот жонглёр ловит на лету тарелки фокусник превратил шляпу в букет цветов а вот клоун он смешит людей у него в куртке петух как интересно в цирке", - "Вот ночная хищная птица голова у нее круглая клюв крючком когти острые Узнали ее это сова она живет в лесах или на чердаках домов ночью птица ловит мышей", - "Черепахи живут на земле и в воде они откладывают яйца прямо на камни черепахи не высиживают их яйца лопаются сами появляются маленькие черепашата а какого размера эти пресмыкающиеся черепахи бывают маленькие и очень большие", - ] - - model_path = str(Path("../models/sbert_punc_case_ru_onnx")) - print(f" ресурсы до старта приложения {gpu_stat(0)}") - time_start = dt.now() - sbertpunc = SbertPuncCaseOnnx(model_path, use_gpu=True) - print(f"Время на инициализацию {(dt.now() - time_start).total_seconds()}") - print(f"Ресурсы после старта приложения {gpu_stat(0)}") - import logging as logger - - # for _ , text_ in enumerate(input_text*100): - # punctuated = asyncio.run(sbertpunc.process(text_) - # ) - # # print(punctuated) - - async def process_texts(texts): - tasks = [sbertpunc.process_punctuation_sessions(text) for text in texts] - results = await asyncio.gather(*tasks, return_exceptions=True) - return results - - - time_start = dt.now() - texts = test_input_text * 100 - punctuated_texts = asyncio.run(process_texts(texts)) - # for text, punctuated in zip(texts, punctuated_texts): - # print(f"Punctuated: {punctuated}\n") - - - - print(f" ресурсы после окончания работы приложения {gpu_stat(0)}") - print(f"Время выполнения {(dt.now() - time_start).total_seconds()}") - -### Увеличение количества сессий в адаптере приводит к пропорциональному увеличению расхода памяти (1,9 * Х для пунктуации) и где-то на 30% быстрее 7,676 против 5,579. -### Если использовать 2 ГПУ, то выполняется где-то на быстрее 60% 4.63946 против 7,676 \ No newline at end of file diff --git a/Recognizer/__init__.py b/Recognizer/__init__.py index f559d56..5819266 100644 --- a/Recognizer/__init__.py +++ b/Recognizer/__init__.py @@ -1,8 +1,8 @@ from starlette.requests import HTTPConnection import multiprocessing import numpy as np +import logging -from utils.do_logging import logger from utils import tokens_to_Result from . import engine from config import settings @@ -10,6 +10,7 @@ import onnx_asr from onnx_asr.loader import PreprocessorRuntimeConfig, OnnxSessionOptions +logger = logging.getLogger(__name__) TENSORRT_providers = ["TensorrtExecutionProvider", "CUDAExecutionProvider", "CPUExecutionProvider"] CUDA_providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] diff --git a/Recognizer/engine/echoe_clearing.py b/Recognizer/engine/echoe_clearing.py index 9c569b0..b124cc1 100644 --- a/Recognizer/engine/echoe_clearing.py +++ b/Recognizer/engine/echoe_clearing.py @@ -1,7 +1,10 @@ # import trash.test_data from typing import List, Dict, Any from difflib import SequenceMatcher -from utils.do_logging import logger +import logging +logger = logging.getLogger(__name__) + + def are_words_similar(word1: str, word2: str, similarity_threshold: float = 0.8) -> bool: """ diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index ebeec24..caa0f74 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -4,6 +4,7 @@ from config import settings import asyncio import uuid +import logging from utils.pre_start_init import ( posted_and_downloaded_audio, audio_buffer, @@ -11,7 +12,7 @@ audio_to_asr, audio_duration, ) -from utils.do_logging import logger + from utils.chunk_doing import find_last_speech_position from utils.resamppling import sync_resample_audiosegment from Recognizer.engine.stream_recognition import simple_recognise, recognise_w_speed_correction, simple_recognise_batch @@ -20,10 +21,12 @@ from Diarisation.diarazer import do_diarizing from threading import Lock +logger = logging.getLogger(__name__) + # Глобальный лок для потокобезопасности audio_lock = Lock() -def process_file(tmp_path, params, recognizer): +def process_file(tmp_path, params, recognizer, punctuator): process_file_start = time.perf_counter() res = False diarized = False @@ -208,9 +211,11 @@ def process_file(tmp_path, params, recognizer): data_to_do_sensitizing = result["diarized_data"] if diarized else result["raw_data"] try: result["sentenced_data"] = asyncio.run(do_sensitizing( - input_asr_json=data_to_do_sensitizing, do_punctuation=params.do_punctuation - ) - ) + input_asr_json=data_to_do_sensitizing, + do_punctuation=params.do_punctuation, + punctuator=punctuator + ) + ) except Exception as e: logger.error(f"do_sensitizing - {e}") error_description = f"do_sensitizing - {e}" diff --git a/Recognizer/engine/sentensizer.py b/Recognizer/engine/sentensizer.py index 42938a9..21b37c3 100644 --- a/Recognizer/engine/sentensizer.py +++ b/Recognizer/engine/sentensizer.py @@ -1,10 +1,10 @@ -from utils.do_logging import logger import numpy as np import asyncio from config import settings -from Punctuation import sbertpunc +import logging +logger = logging.getLogger(__name__) -async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): +async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False, punctuator=None): """ :param do_punctuation: Если True, то производит пунктуацию и капитализацию над собранными в предложения выражения. :param input_asr_json: {"channel_{n_channel + 1}": @@ -64,7 +64,7 @@ async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): # Обработка случая с одним словом text = words[0].get('word') if do_punctuation: - text = await sbertpunc.punctuate(text) + text = await punctuator.punctuate(text) sentence_element.append({ "start": start_time, @@ -108,7 +108,7 @@ async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): else: if do_punctuation: - text = await sbertpunc.punctuate(' '.join(str(word) for word in sentences)) + text = await punctuator.punctuate(' '.join(str(word) for word in sentences)) else: text = ' '.join(str(word) for word in sentences) @@ -128,7 +128,7 @@ async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): if do_punctuation: - text = await sbertpunc.punctuate(' '.join(str(word) for word in sentences)) + text = await punctuator.punctuate(' '.join(str(word) for word in sentences)) else: text = ' '.join(str(word) for word in sentences) @@ -143,7 +143,7 @@ async def do_sensitizing(input_asr_json: str, do_punctuation: bool = False): if do_punctuation: # Тут он текст разделит сам. - one_text_only = str(await sbertpunc.punctuate(one_text_only)) + one_text_only = str(await punctuator.punctuate(one_text_only)) text_only.append(one_text_only) diff --git a/Recognizer/engine/stream_recognition.py b/Recognizer/engine/stream_recognition.py index 223f668..fda8780 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -4,11 +4,12 @@ from utils.bytes_to_samples_audio import get_np_array_samples_float32 from utils.resamppling import sync_resample_audiosegment from utils.slow_down_audio import do_slow_down_audio -from utils.do_logging import logger +import logging from utils.chunk_doing import samples_padding from dataclasses import asdict from onnx_asr.utils import read_wav_files, pad_list +logger = logging.getLogger(__name__) def calc_speed(data): time_to_speak_tokens = 0 diff --git a/VoiceActivityDetector/__init__.py b/VoiceActivityDetector/__init__.py index 7dc52c8..dfe071b 100644 --- a/VoiceActivityDetector/__init__.py +++ b/VoiceActivityDetector/__init__.py @@ -1,8 +1,8 @@ from config import settings from utils.pre_start_init import paths -from utils.do_logging import logger import requests - +import logging +logger = logging.getLogger(__name__) if not paths.get("vad_model_path").exists(): logger.info("Модель silero_vad.onnx отсутствует. Предпринимаем попытку скачать её.") diff --git a/VoiceActivityDetector/do_vad.py b/VoiceActivityDetector/do_vad.py index a8777ea..eee0b34 100644 --- a/VoiceActivityDetector/do_vad.py +++ b/VoiceActivityDetector/do_vad.py @@ -1,12 +1,13 @@ import datetime import asyncio - +import logging import numpy as np import onnxruntime as ort from pydub import AudioSegment from pathlib import Path from utils.resamppling import sync_resample_audiosegment -from utils.do_logging import logger + +logger = logging.getLogger(__name__) class SileroVAD: diff --git a/main.py b/main.py index 4f06187..1557367 100644 --- a/main.py +++ b/main.py @@ -1,4 +1,3 @@ -from Recognizer import Recognizer import logging import uvicorn from config import settings @@ -58,8 +57,15 @@ async def lifespan(app): # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) + # Инициируем recognizer + from Recognizer import Recognizer app.state.recognizer = Recognizer() + # Инициируем punctuator + from Punctuation import SbertPuncCaseOnnx + app.state.punctuator = SbertPuncCaseOnnx(paths.get("punctuation_model_path"),use_gpu=settings.PUNCTUATE_WITH_GPU) + + if settings.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( target=lambda: start_file_watcher(file_path=str(paths.get("local_recognition_folder"))), @@ -69,8 +75,12 @@ async def lifespan(app): logger.info("File watcher started") yield # Здесь приложение работает + # cleanup (если нужно) - del app.state.recognize + if hasattr(app.state, "recognizer"): + del app.state.recognizer + if hasattr(app.state, "punctuator"): + del app.state.punctuator app = FastAPI( diff --git a/routes/legacy/demo_page.py b/routes/legacy/demo_page.py index 7b13da8..ee5b983 100644 --- a/routes/legacy/demo_page.py +++ b/routes/legacy/demo_page.py @@ -1,9 +1,9 @@ from fastapi import APIRouter, Request -from utils.do_logging import logger from fastapi.responses import HTMLResponse from fastapi.templating import Jinja2Templates - +import logging +logger = logging.getLogger(__name__) router = APIRouter() diff --git a/routes/legacy/post_by_file_FORM.py b/routes/legacy/post_by_file_FORM.py index f53cdcd..88239ed 100644 --- a/routes/legacy/post_by_file_FORM.py +++ b/routes/legacy/post_by_file_FORM.py @@ -1,15 +1,15 @@ from io import BytesIO import asyncio - +import logging from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile -from utils.do_logging import logger from models.fast_api_models import PostFileRequest, BaseResponse from Recognizer import get_recognizer, Recognizer from Recognizer.engine.file_recognition import process_file -from threading import Lock - +from Punctuation import get_punctuator, SbertPuncCaseOnnx +import logging +logger = logging.getLogger(__name__) router = APIRouter() # Функция для извлечения параметров из FormData @@ -45,7 +45,8 @@ def get_file_request( async def async_receive_file_legacy( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), - recognizer: Recognizer = Depends(get_recognizer) + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) ) -> BaseResponse: # Сохраняем файл на диск асинхронно try: @@ -65,7 +66,11 @@ async def async_receive_file_legacy( logger.info(f"Получен и сохранён файл {file.filename}") try: # Запускаем обработку в потоке - result_dict = await asyncio.to_thread(process_file, buffer, params, recognizer) + result_dict = await asyncio.to_thread(process_file, + tmp_path=buffer, + params=params, + recognizer=recognizer, + punctuator=punctuator) result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/legacy/post_by_url.py b/routes/legacy/post_by_url.py index 959467e..341f437 100644 --- a/routes/legacy/post_by_url.py +++ b/routes/legacy/post_by_url.py @@ -3,13 +3,14 @@ import os from fastapi import APIRouter from utils.pre_start_init import posted_and_downloaded_audio -from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, BaseResponse from fastapi import Depends from Recognizer import get_recognizer, Recognizer +from Punctuation import get_punctuator, SbertPuncCaseOnnx + from Recognizer.engine.file_recognition import process_file from threading import Lock from io import BytesIO @@ -22,7 +23,8 @@ @router.post("/post_one_step_req", response_model=BaseResponse) async def post(params: SyncASRRequest, - recognizer: Recognizer = Depends(get_recognizer) + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) ) -> BaseResponse: """ На вход ждёт str(HttpUrl) - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n @@ -60,7 +62,11 @@ async def post(params: SyncASRRequest, try: # Запускаем обработку в потоке - result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params, recognizer) + result_dict = await asyncio.to_thread(process_file, + tmp_path =posted_and_downloaded_audio[post_id], + params=params, + recognizer=recognizer, + punctuator=punctuator) result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 3d980b0..2e0d7e9 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -1,11 +1,14 @@ from io import BytesIO import asyncio -from utils.do_logging import logger from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile from models.fast_api_models import PostFileRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file from Recognizer import get_recognizer, Recognizer +from Punctuation import get_punctuator, SbertPuncCaseOnnx + +import logging +logger = logging.getLogger(__name__) router = APIRouter() # Функция для извлечения параметров из FormData @@ -41,7 +44,9 @@ def get_file_request( async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), - recognizer: Recognizer = Depends(get_recognizer) + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) + ) -> V1ASRResponse: try: buffer = BytesIO(await file.read()) @@ -57,7 +62,11 @@ async def async_receive_file( else: logger.info(f"Получен и сохранён файл {file.filename}") try: - result_dict = await asyncio.to_thread(process_file, buffer, params, recognizer) + result_dict = await asyncio.to_thread(process_file, + tmp_path=buffer, + params=params, + recognizer=recognizer, + punctuator=punctuator) return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/routes/v1/post_by_url.py b/routes/v1/post_by_url.py index b55c3d4..f43fa72 100644 --- a/routes/v1/post_by_url.py +++ b/routes/v1/post_by_url.py @@ -1,22 +1,25 @@ import uuid import asyncio -from fastapi import APIRouter from utils.pre_start_init import posted_and_downloaded_audio -from utils.do_logging import logger from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData -from fastapi import Depends +from fastapi import APIRouter, Depends from Recognizer import get_recognizer, Recognizer +from Punctuation import get_punctuator, SbertPuncCaseOnnx from Recognizer.engine.file_recognition import process_file +import logging +logger = logging.getLogger(__name__) router = APIRouter() @router.post("/post_one_step_req", response_model=V1ASRResponse) async def post_v1(params: SyncASRRequest, - recognizer: Recognizer = Depends(get_recognizer)) -> V1ASRResponse: + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) + ) -> V1ASRResponse: post_id = uuid.uuid4() if params.AudioFileUrl: res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 685f224..891e9c3 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -1,12 +1,12 @@ from pydub import AudioSegment import ujson +import logging from config import settings import uuid from io import BytesIO from fastapi import APIRouter, WebSocket -from utils.do_logging import logger from utils.chunk_doing import find_last_speech_position from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration,ws_collected_asr_res from utils.send_messages import send_messages @@ -18,13 +18,17 @@ from Recognizer.engine.sentensizer import do_sensitizing from Recognizer.engine.stream_recognition import simple_recognise - +from Punctuation import get_punctuator, SbertPuncCaseOnnx router = APIRouter() +logger = logging.getLogger(__name__) + @router.websocket("/ws") async def websocket(ws: WebSocket, - recognizer: Recognizer = Depends(get_recognizer)): + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + ): wait_null_answers=True client_id = uuid.uuid4() logger.debug(f'Принят новый сокет id = {client_id}') @@ -218,7 +222,7 @@ async def websocket(ws: WebSocket, if do_dialogue: try: - sentenced_data = await do_sensitizing(ws_collected_asr_res[client_id], do_punctuation) + sentenced_data = await do_sensitizing(ws_collected_asr_res[client_id], do_punctuation, punctuator=punctuator) except Exception as e: logger.error(f"await do_sensitizing - {e}") error_description = f"do_sensitizing - {e}" diff --git a/utils/chunk_doing.py b/utils/chunk_doing.py index 3d81f92..cbb1c96 100644 --- a/utils/chunk_doing.py +++ b/utils/chunk_doing.py @@ -1,7 +1,6 @@ from config import settings import numpy as np from pydub import AudioSegment -from utils.do_logging import logger from utils.bytes_to_samples_audio import get_np_array_samples_float32 from utils.pre_start_init import (audio_overlap, @@ -11,6 +10,8 @@ from VoiceActivityDetector import vad from utils.resamppling import async_resample_audiosegment +import logging +logger = logging.getLogger(__name__) async def find_last_speech_position(socket_id, is_last_chunk): diff --git a/utils/file_exists.py b/utils/file_exists.py index 9e07e19..d9ae4ca 100644 --- a/utils/file_exists.py +++ b/utils/file_exists.py @@ -1,5 +1,6 @@ from pathlib import Path -from utils.do_logging import logger +import logging +logger = logging.getLogger(__name__) def assert_file_exists(filename_path: Path): diff --git a/utils/save_asr_to_md.py b/utils/save_asr_to_md.py index bb55286..fe5db49 100644 --- a/utils/save_asr_to_md.py +++ b/utils/save_asr_to_md.py @@ -1,14 +1,16 @@ import json +import logging from datetime import datetime -from utils.do_logging import logger from pathlib import Path -import config +from config import settings + +logger = logging.getLogger(__name__) async def save_to_file(json_data: json, file_name: Path = None): from utils.pre_start_init import paths - file_extension = "md" if config.HUMAN_FORMAT_MD_FILE else "json" + file_extension = "md" if settings.HUMAN_FORMAT_MD_FILE else "json" output_folder = paths.get("result_local_recognition_folder") local_path = paths.get("local_recognition_folder") # Получаем корневую папку для отслеживания @@ -34,7 +36,7 @@ async def save_to_file(json_data: json, file_name: Path = None): file_path.parent.mkdir(parents=True, exist_ok=True) # С - if config.HUMAN_FORMAT_MD_FILE: + if settings.HUMAN_FORMAT_MD_FILE: with open(file_path, "w", encoding="utf-8") as f: f.write(f"**Дата:** {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n") try: diff --git a/utils/send_messages.py b/utils/send_messages.py index 4ac6697..0ea51e3 100644 --- a/utils/send_messages.py +++ b/utils/send_messages.py @@ -1,5 +1,5 @@ -from utils.do_logging import logger - +import logging +logger = logging.getLogger(__name__) async def send_messages(_socket, _channel_name = None, _data=None, _silence=True, _error=None, _last_message=False, _sentenced_data=None): diff --git a/utils/slow_down_audio.py b/utils/slow_down_audio.py index f227803..afc4a91 100644 --- a/utils/slow_down_audio.py +++ b/utils/slow_down_audio.py @@ -1,9 +1,9 @@ import numpy as np from pydub import AudioSegment from utils.bytes_to_samples_audio import get_np_array_samples_float32 -from utils.do_logging import logger from scipy.interpolate import CubicSpline -from pathlib import Path +import logging +logger = logging.getLogger(__name__) async def do_slow_down_audio(audio_segment: AudioSegment, slowdown_rate: float) -> AudioSegment: """ diff --git a/utils/tokens_to_Result.py b/utils/tokens_to_Result.py index 9dcb52b..2b4f00d 100644 --- a/utils/tokens_to_Result.py +++ b/utils/tokens_to_Result.py @@ -1,9 +1,5 @@ -import asyncio - -from numpy.ma.core import count -from sympy.physics.units import speed - -from utils.do_logging import logger +import logging +logger = logging.getLogger(__name__) # Парсим JSON def process_multi_tokens_vocab_output(input_json, time_shift = 0.0, multiplier=1): From c6752558ed4cefa82af0fbab866835e143337133 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 30 Apr 2026 16:45:41 +0300 Subject: [PATCH 17/46] =?UTF-8?q?=D0=B8=D1=81=D0=BF=D1=80=D0=B0=D0=B2?= =?UTF-8?q?=D0=BB=D0=B5=D0=BD=D0=B8=D0=B5=20=D0=BE=D1=88=D0=B8=D0=B1=D0=BA?= =?UTF-8?q?=D0=B8=20=D0=BD=D0=B5=D0=BF=D1=80=D0=B0=D0=B2=D0=B8=D0=BB=D1=8C?= =?UTF-8?q?=D0=BD=D0=BE=D0=B9=20=D0=BE=D0=B1=D1=80=D0=B0=D0=B1=D0=BE=D1=82?= =?UTF-8?q?=D0=BA=D0=B8=20"=D0=BF=D0=BE=D1=81=D0=BB=D0=B5=D0=B4=D0=BD?= =?UTF-8?q?=D0=B5=D0=B3=D0=BE=20=D0=BD=D0=B5=D0=BF=D0=BE=D0=BB=D0=BD=D0=BE?= =?UTF-8?q?=D0=B3=D0=BE=20=D1=87=D0=B0=D0=BD=D0=BA=D0=B0"=20=D0=B2=20/ws?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- routes/ws_audio_transkrib.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 891e9c3..bece710 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -85,6 +85,10 @@ async def websocket(ws: WebSocket, chunk = message.get('bytes') if audio_format == 'pcm16': + # Проверяем и добавляем недостающие нулевые байты в чанки. + if len(chunk) % 2 != 0: + chunk += bytes(2 - (len(chunk) % 2)) + # Переводим чанк в объект Audiosegment audiosegment_chunk = AudioSegment( chunk, @@ -92,6 +96,7 @@ async def websocket(ws: WebSocket, sample_width = 2, # Ширина сэмпла (2 байта для int16) channels = 1 # Количество каналов. По умолчанию - 1, Моно. ) + else: try: buffer = BytesIO(chunk) @@ -110,7 +115,6 @@ async def websocket(ws: WebSocket, if audiosegment_chunk.channels != 1: audiosegment_chunk = audiosegment_chunk.set_channels(1) - # Копим буфер audio_buffer[client_id] += audiosegment_chunk @@ -118,6 +122,7 @@ async def websocket(ws: WebSocket, if (audio_overlap[client_id]+audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: # Проверяем новый чанк перед объединением (там же режем хвост и добавляем его при необходимости) await find_last_speech_position(client_id, is_last_chunk=False) + else: continue except Exception as e: From 746e336549a1bfb27ef7564877e5c283813d451b Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 30 Apr 2026 17:37:05 +0300 Subject: [PATCH 18/46] =?UTF-8?q?=D0=9F=D0=B5=D1=80=D0=B5=D0=BD=D0=BE?= =?UTF-8?q?=D1=81=20=D0=B8=D0=BD=D0=B8=D1=86=D0=B8=D0=B0=D0=BB=D0=B8=D0=B7?= =?UTF-8?q?=D0=B0=D1=86=D0=B8=D0=B8=20diarizer=20=D0=B2=20main.py?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Diarisation/__init__.py | 41 ++++++++++----------------- Diarisation/diarazer.py | 8 ++---- Diarisation/do_diarize.py | 1 + Recognizer/engine/file_recognition.py | 7 +++-- main.py | 20 +++++++++++++ routes/legacy/post_by_file_FORM.py | 9 ++++-- routes/legacy/post_by_url.py | 24 +++++++--------- routes/v1/post_by_file_FORM.py | 15 +++++++--- routes/v1/post_by_url.py | 14 +++++++-- 9 files changed, 83 insertions(+), 56 deletions(-) diff --git a/Diarisation/__init__.py b/Diarisation/__init__.py index 79f12dc..bab5650 100644 --- a/Diarisation/__init__.py +++ b/Diarisation/__init__.py @@ -3,25 +3,27 @@ from VoiceActivityDetector import vad import requests import logging +from starlette.requests import HTTPConnection logger = logging.getLogger(__name__) -if settings.CAN_DIAR: + +def ensure_diar_model() -> bool: + """Проверяет наличие модели диаризации и при необходимости скачивает её.""" + if not settings.CAN_DIAR: + return False if not paths.get("diar_speaker_model_path").exists(): logger.error(f"Модель для диаризации не найдена. Предпринимаются попытки скачать {settings.DIAR_MODEL_NAME}") output_path = paths.get("diar_speaker_model_path") api_url = "https://modelscope.cn/api/v1/datasets/wenet/wespeaker_pretrained_models/oss/tree" - try: # Предпринимаем попытки скачать и положить файл в нужное место. - # 1. Получаем JSON с данными файлов + try: response = requests.get(api_url, headers={"User-Agent": "Mozilla/5.0"}) response.raise_for_status() - # 2. Ищем нужный файл target_file = settings.DIAR_MODEL_NAME file_data = next((item for item in response.json()["Data"] if item["Key"] == target_file), None) if not file_data: logger.error(f"Модели с именем {settings.DIAR_MODEL_NAME} в списке возможных для загрузки нет.") - # Получаем список всех ONNX-моделей onnx_models_with_size = [ (item['Key'].split(".")[0], item['Size'] // (1024 * 1024)) for item in response.json()["Data"] @@ -30,10 +32,8 @@ logger.info("Доступны для скачивания следующие модели:") for name, size in sorted(onnx_models_with_size): logger.info(f"{name} - {size} MB") - raise Exception("Файл не найден в API") - # 3. Скачиваем по прямой ссылке download_url = file_data["Url"] with requests.get(download_url, stream=True) as r: r.raise_for_status() @@ -41,31 +41,20 @@ for chunk in r.iter_content(chunk_size=8192): f.write(chunk) except Exception as e: - logger.error(f"Файл модели не сохранён.Ошибка - {e})") - + logger.error(f"Файл модели не сохранён. Ошибка - {e}") + return False else: logger.info(f"Модель успешно загружена : {output_path}") - else: logger.debug(f"Будет использован имеющийся файл {settings.DIAR_MODEL_NAME}") - if not paths.get("diar_speaker_model_path").exists(): - logger.error(f"Модель для Диаризации отсутствует. Диризация выключена и будет не доступна.") + logger.error(f"Модель для Диаризации отсутствует. Диаризация выключена.") logger.error(f"Скачайте модель со страницы 'https://github.com/wenet-e2e/wespeaker/blob/master/docs/pretrained.md' " - f"и поместите в по адресу: {str(paths.get('diar_speaker_model_path'))}") - settings.CAN_DIAR = False - else: - from .do_diarize import Diarizer + f"и поместите по адресу: {str(paths.get('diar_speaker_model_path'))}") + return False + return True - diarizer = Diarizer(embedding_model_path=paths.get("diar_speaker_model_path"), - vad=vad, - max_phrase_gap=1, - batch_size=settings.DIAR_GPU_BATCH_SIZE, - cpu_workers=settings.CPU_WORKERS, - use_gpu=settings.DIAR_WITH_GPU - ) - logger.info(f"Успешно загружена модель Диаризации") -else: - logger.info(f"Диаризация не включена и будет недоступна.") \ No newline at end of file +def get_diarizer(conn: HTTPConnection): + return conn.app.state.diarizer diff --git a/Diarisation/diarazer.py b/Diarisation/diarazer.py index 091446b..a1b825e 100644 --- a/Diarisation/diarazer.py +++ b/Diarisation/diarazer.py @@ -1,19 +1,15 @@ import datetime - -from config import settings from Diarisation.do_diarize import load_and_preprocess_audio from utils.pre_start_init import posted_and_downloaded_audio from collections import defaultdict import logging logger = logging.getLogger(__name__) -if settings.CAN_DIAR: - from Diarisation import diarizer - async def do_diarizing( file_id:str, asr_raw_data, + diarizer, num_speakers:int = -1, filter_cutoff:int = 50, filter_order:int = 10, @@ -220,4 +216,4 @@ def group_words(words): # min_overlap_ratio=0.5, # max_pause=1.5, # group_unmatched_words=True # Включить группировку несопоставленных слов -# ) ) \ No newline at end of file +# ) ) diff --git a/Diarisation/do_diarize.py b/Diarisation/do_diarize.py index ec7b821..b4d71fa 100644 --- a/Diarisation/do_diarize.py +++ b/Diarisation/do_diarize.py @@ -17,6 +17,7 @@ import logging logger = logging.getLogger(__name__) + class Diarizer: def __init__(self, embedding_model_path: str, vad, diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index caa0f74..fc8a84f 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -26,7 +26,7 @@ # Глобальный лок для потокобезопасности audio_lock = Lock() -def process_file(tmp_path, params, recognizer, punctuator): +def process_file(tmp_path, params, recognizer, punctuator, diarizer): process_file_start = time.perf_counter() res = False diarized = False @@ -198,7 +198,10 @@ def process_file(tmp_path, params, recognizer, punctuator): if params.do_diarization: try: result["diarized_data"] = asyncio.run(do_diarizing( - file_id=str(post_id), asr_raw_data=result["raw_data"], diar_vad_sensity=params.diar_vad_sensity + file_id=str(post_id), + asr_raw_data=result["raw_data"], + diar_vad_sensity=params.diar_vad_sensity, + diarizer=diarizer, )) except Exception as e: logger.error(f"do_diarizing - {e}") diff --git a/main.py b/main.py index 1557367..072681d 100644 --- a/main.py +++ b/main.py @@ -19,6 +19,7 @@ from utils.files_whatcher import start_file_watcher from utils.pre_start_init import paths import threading +from VoiceActivityDetector import vad from routes.ws_audio_transkrib import router as ws_audio_transkrib_router from routes.v1 import router as v1_router @@ -65,6 +66,23 @@ async def lifespan(app): from Punctuation import SbertPuncCaseOnnx app.state.punctuator = SbertPuncCaseOnnx(paths.get("punctuation_model_path"),use_gpu=settings.PUNCTUATE_WITH_GPU) + # Инициируем diarizer + if settings.CAN_DIAR: + from Diarisation import ensure_diar_model + if ensure_diar_model(): + from Diarisation.do_diarize import Diarizer + app.state.diarizer = Diarizer( + embedding_model_path=paths.get("diar_speaker_model_path"), + vad=vad, # todo- моежт быть использовать разные VAD для диаризации и разделения на чанки? + max_phrase_gap=1, + batch_size=settings.DIAR_GPU_BATCH_SIZE, + cpu_workers=settings.CPU_WORKERS, + use_gpu=settings.DIAR_WITH_GPU + ) + logger.info("Модель диаризации загружена") + else: + settings.CAN_DIAR = False + logger.warning("Диаризация недоступна: модель не найдена и не удалось скачать") if settings.DO_LOCAL_FILE_RECOGNITIONS: observer_thread = threading.Thread( @@ -81,6 +99,8 @@ async def lifespan(app): del app.state.recognizer if hasattr(app.state, "punctuator"): del app.state.punctuator + if hasattr(app.state, "diarizer"): + del app.state.diarizer app = FastAPI( diff --git a/routes/legacy/post_by_file_FORM.py b/routes/legacy/post_by_file_FORM.py index 88239ed..83b2c61 100644 --- a/routes/legacy/post_by_file_FORM.py +++ b/routes/legacy/post_by_file_FORM.py @@ -8,6 +8,9 @@ from Recognizer.engine.file_recognition import process_file from Punctuation import get_punctuator, SbertPuncCaseOnnx +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer + import logging logger = logging.getLogger(__name__) router = APIRouter() @@ -46,7 +49,8 @@ async def async_receive_file_legacy( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) ) -> BaseResponse: # Сохраняем файл на диск асинхронно try: @@ -70,7 +74,8 @@ async def async_receive_file_legacy( tmp_path=buffer, params=params, recognizer=recognizer, - punctuator=punctuator) + punctuator=punctuator, + diarizer=diarizer) result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/legacy/post_by_url.py b/routes/legacy/post_by_url.py index 341f437..472cd8b 100644 --- a/routes/legacy/post_by_url.py +++ b/routes/legacy/post_by_url.py @@ -1,30 +1,26 @@ import uuid import asyncio -import os -from fastapi import APIRouter +from fastapi import APIRouter, Depends from utils.pre_start_init import posted_and_downloaded_audio from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, BaseResponse -from fastapi import Depends from Recognizer import get_recognizer, Recognizer - from Punctuation import get_punctuator, SbertPuncCaseOnnx - from Recognizer.engine.file_recognition import process_file -from threading import Lock -from io import BytesIO +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer +import logging +logger = logging.getLogger(__name__) router = APIRouter() -# Глобальный лок для потокобезопасности -audio_lock = Lock() - @router.post("/post_one_step_req", response_model=BaseResponse) async def post(params: SyncASRRequest, recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) ) -> BaseResponse: """ На вход ждёт str(HttpUrl) - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n @@ -63,10 +59,12 @@ async def post(params: SyncASRRequest, try: # Запускаем обработку в потоке result_dict = await asyncio.to_thread(process_file, - tmp_path =posted_and_downloaded_audio[post_id], + tmp_path=posted_and_downloaded_audio[post_id], params=params, recognizer=recognizer, - punctuator=punctuator) + punctuator=punctuator, + diarizer=diarizer) + result = BaseResponse(**result_dict) except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" diff --git a/routes/v1/post_by_file_FORM.py b/routes/v1/post_by_file_FORM.py index 2e0d7e9..413b75d 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/routes/v1/post_by_file_FORM.py @@ -7,10 +7,14 @@ from Recognizer import get_recognizer, Recognizer from Punctuation import get_punctuator, SbertPuncCaseOnnx +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer + import logging -logger = logging.getLogger(__name__) +logger = logging.getLogger(__name__) router = APIRouter() + # Функция для извлечения параметров из FormData def get_file_request( keep_raw: bool = Form(default=True, description="Сохранять сырые данные."), @@ -45,9 +49,11 @@ async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) - + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) ) -> V1ASRResponse: + + try: buffer = BytesIO(await file.read()) buffer.seek(0) @@ -66,7 +72,8 @@ async def async_receive_file( tmp_path=buffer, params=params, recognizer=recognizer, - punctuator=punctuator) + punctuator=punctuator, + diarizer=diarizer) return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/routes/v1/post_by_url.py b/routes/v1/post_by_url.py index f43fa72..c89b864 100644 --- a/routes/v1/post_by_url.py +++ b/routes/v1/post_by_url.py @@ -6,9 +6,11 @@ from fastapi import APIRouter, Depends from Recognizer import get_recognizer, Recognizer +from Recognizer.engine.file_recognition import process_file from Punctuation import get_punctuator, SbertPuncCaseOnnx +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer -from Recognizer.engine.file_recognition import process_file import logging logger = logging.getLogger(__name__) @@ -18,7 +20,8 @@ @router.post("/post_one_step_req", response_model=V1ASRResponse) async def post_v1(params: SyncASRRequest, recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator) + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) ) -> V1ASRResponse: post_id = uuid.uuid4() if params.AudioFileUrl: @@ -35,7 +38,12 @@ async def post_v1(params: SyncASRRequest, ) try: - result_dict = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params, recognizer) + result_dict = await asyncio.to_thread(process_file, + tmp_path=posted_and_downloaded_audio[post_id], + params=params, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer) return V1ASRResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), From 203cadb7adaffc7a484bf5e27582f795bceb39ab Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 4 May 2026 14:18:38 +0300 Subject: [PATCH 19/46] =?UTF-8?q?=D0=9F=D0=BE=D0=B4=D0=B3=D0=BE=D1=82?= =?UTF-8?q?=D0=BE=D0=B2=D0=BA=D0=B0=20=D0=BA=20=D0=B0=D0=B2=D1=82=D0=BE?= =?UTF-8?q?=D1=80=D0=B8=D0=B7=D0=B0=D1=86=D0=B8=D0=B8.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/engine/file_recognition.py | 4 +- config.py | 25 +++++- core/exception_handlers.py | 51 +++++++++++++ core/exceptions.py | 19 +++++ core/security.py | 84 +++++++++++++++++++++ main.py | 72 ++++++------------ models/fast_api_models.py | 16 ---- requirements.txt | 6 +- tests/test_gzip_middleware.py | 39 +++++++++- tests/test_security.py | 105 ++++++++++++++++++++++++++ 10 files changed, 351 insertions(+), 70 deletions(-) create mode 100644 core/exception_handlers.py create mode 100644 core/exceptions.py create mode 100644 core/security.py create mode 100644 tests/test_security.py diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index fc8a84f..2b15457 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -69,14 +69,14 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): return result # Приводим фреймрейт к фреймрейту модели - print(f"Начало проверки фреймрейта {(time.perf_counter()-process_file_start):.4f} сек.") + logger.debug(f"Начало проверки фреймрейта {(time.perf_counter()-process_file_start):.4f} сек.") try: with audio_lock: if posted_and_downloaded_audio[post_id].frame_rate != settings.BASE_SAMPLE_RATE: posted_and_downloaded_audio[post_id] = sync_resample_audiosegment( audio_data=posted_and_downloaded_audio[post_id], target_sample_rate=settings.BASE_SAMPLE_RATE) - print(f"Корректировка фреймрейта {(time.perf_counter() - process_file_start):.4f} сек.") + logger.debug(f"Корректировка фреймрейта {(time.perf_counter() - process_file_start):.4f} сек.") except KeyError as e_key: error_description = f"Ошибка обращения по ключу {post_id} при изменения фреймрейта - {e_key}" logger.error(error_description) diff --git a/config.py b/config.py index 0bce7a4..cfdcf59 100644 --- a/config.py +++ b/config.py @@ -1,6 +1,6 @@ import logging import os -from datetime import date +from datetime import date, timedelta from pydantic_settings import BaseSettings, SettingsConfigDict from pydantic import Field, field_validator, model_validator @@ -29,7 +29,7 @@ class Settings(BaseSettings): MODEL_NAME: str = "gigaam-v3-ctc" # Стрим из астериска отдаёт только 8к BASE_SAMPLE_RATE: int = 16000 - PROVIDER: str = "CUDA" + PROVIDER: str = "CPU" NUM_THREADS: int = 0 # HuggingFace Hub settings @@ -104,6 +104,13 @@ class Settings(BaseSettings): DELETE_LOCAL_FILE_AFTR_ASR: bool = False HUMAN_FORMAT_MD_FILE: bool = False + # Auth / JWT settings + SECRET_KEY: str = Field(default="change-me-in-production-32-chars-long", min_length=32) + ALGORITHM: str = "HS256" + ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 + REFRESH_TOKEN_EXPIRE_MINUTES: int = 10080 # 7 дней + BCRYPT_ROUNDS: int = 12 + @field_validator( 'IS_PROD', 'MAKE_MONO', 'USE_BATCH', 'VAD_WITH_GPU', 'CAN_PUNCTUATE', 'PUNCTUATE_WITH_GPU', 'CAN_DIAR', @@ -151,6 +158,20 @@ def _compute_derived(self): "Please specify explicit origins in CORS_ORIGINS." ) + # Валидация SECRET_KEY для production + if self.IS_PROD and len(self.SECRET_KEY) < 32: + raise ValueError( + "SECURITY ERROR: SECRET_KEY must be at least 32 characters long in production (IS_PROD=True). " + "Please set a strong SECRET_KEY environment variable." + ) + + # Warning при использовании дефолтного SECRET_KEY в production + if self.IS_PROD and self.SECRET_KEY == "change-me-in-production-32-chars-long": + logging.warning( + "SECURITY WARNING: Using default SECRET_KEY in production (IS_PROD=True). " + "Please set a strong unique SECRET_KEY environment variable." + ) + return self diff --git a/core/exception_handlers.py b/core/exception_handlers.py new file mode 100644 index 0000000..f1d1f41 --- /dev/null +++ b/core/exception_handlers.py @@ -0,0 +1,51 @@ +import logging +import traceback + +from fastapi import Request +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse +from starlette.exceptions import HTTPException as StarletteHTTPException + +from models.fast_api_models import ErrorResponse + +logger = logging.getLogger(__name__) + + +async def validation_exception_handler(request: Request, exc: RequestValidationError): + return JSONResponse( + status_code=422, + content=ErrorResponse( + success=False, + error_description="Validation error", + details=str(exc), + ).model_dump() + ) + + +async def http_exception_handler(request: Request, exc: StarletteHTTPException): + return JSONResponse( + status_code=exc.status_code, + content=ErrorResponse( + success=False, + error_description=exc.detail, + ).model_dump() + ) + + +async def general_exception_handler(request: Request, exc: Exception): + logger.error(f"Unhandled exception: {exc}\n{traceback.format_exc()}") + return JSONResponse( + status_code=500, + content=ErrorResponse( + success=False, + error_description="Internal server error", + details=str(exc), + ).model_dump() + ) + + +def register_exception_handlers(app): + """Регистрация обработчиков исключений на экземпляре FastAPI.""" + app.add_exception_handler(RequestValidationError, validation_exception_handler) + app.add_exception_handler(StarletteHTTPException, http_exception_handler) + app.add_exception_handler(Exception, general_exception_handler) diff --git a/core/exceptions.py b/core/exceptions.py new file mode 100644 index 0000000..dfdbb60 --- /dev/null +++ b/core/exceptions.py @@ -0,0 +1,19 @@ +from fastapi import HTTPException + + +class CredentialsException(HTTPException): + """401 — не удалось проверить учётные данные.""" + def __init__(self): + super().__init__(status_code=401, detail="Could not validate credentials") + + +class TokenExpiredException(HTTPException): + """401 — токен просрочен.""" + def __init__(self): + super().__init__(status_code=401, detail="Token has expired") + + +class InvalidTokenException(HTTPException): + """401 — невалидный токен.""" + def __init__(self): + super().__init__(status_code=401, detail="Invalid token") diff --git a/core/security.py b/core/security.py new file mode 100644 index 0000000..fddc796 --- /dev/null +++ b/core/security.py @@ -0,0 +1,84 @@ +import bcrypt +import jwt + +from datetime import datetime, timedelta, timezone +from typing import Literal + +from pydantic import BaseModel + +from config import settings +from core.exceptions import InvalidTokenException, TokenExpiredException + + +class TokenPayload(BaseModel): + """Типизированная модель payload JWT-токена.""" + sub: str + exp: datetime + iat: datetime + type: Literal["access", "refresh"] + + +def get_password_hash(password: str) -> str: + """Генерация хеша пароля через bcrypt с использованием настроек из конфига.""" + salt = bcrypt.gensalt(rounds=settings.BCRYPT_ROUNDS) + hashed = bcrypt.hashpw(password.encode("utf-8"), salt) + return hashed.decode("utf-8") + + +def verify_password(plain_password: str, hashed_password: str) -> bool: + """Проверка пароля против хеша через bcrypt.""" + return bcrypt.checkpw( + plain_password.encode("utf-8"), + hashed_password.encode("utf-8"), + ) + + +def create_access_token(data: dict, expires_delta: timedelta | None = None) -> str: + """Создание access-токена.""" + to_encode = data.copy() + now = datetime.now(timezone.utc) + if expires_delta: + expire = now + expires_delta + else: + expire = now + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) + to_encode.update({"exp": expire, "iat": now, "type": "access"}) + return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) + + +def create_refresh_token(data: dict, expires_delta: timedelta | None = None) -> str: + """Создание refresh-токена.""" + to_encode = data.copy() + now = datetime.now(timezone.utc) + if expires_delta: + expire = now + expires_delta + else: + expire = now + timedelta(minutes=settings.REFRESH_TOKEN_EXPIRE_MINUTES) + to_encode.update({"exp": expire, "iat": now, "type": "refresh"}) + return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) + + +def decode_token(token: str, expected_type: Literal["access", "refresh"] | None = None) -> TokenPayload: + """Декодирование и валидация JWT-токена.""" + try: + payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) + except jwt.ExpiredSignatureError: + raise TokenExpiredException() + except jwt.InvalidTokenError: + raise InvalidTokenException() + + token_type = payload.get("type") + if expected_type is not None and token_type != expected_type: + raise InvalidTokenException() + + exp = payload.get("exp") + iat = payload.get("iat") + if exp is None: + raise InvalidTokenException() + exp_dt = datetime.fromtimestamp(exp, tz=timezone.utc) + iat_dt = datetime.fromtimestamp(iat, tz=timezone.utc) if iat is not None else exp_dt + + sub = payload.get("sub") + if sub is None: + raise InvalidTokenException() + + return TokenPayload(sub=sub, exp=exp_dt, iat=iat_dt, type=token_type) diff --git a/main.py b/main.py index 072681d..a53d6d5 100644 --- a/main.py +++ b/main.py @@ -5,17 +5,15 @@ import gc from contextlib import asynccontextmanager from fastapi import FastAPI, Request -from fastapi.exceptions import RequestValidationError -from fastapi.responses import JSONResponse -from starlette.exceptions import HTTPException as StarletteHTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles -from starlette.middleware.base import BaseHTTPMiddleware +from starlette.datastructures import MutableHeaders from starlette.middleware.gzip import GZipMiddleware from starlette.middleware.trustedhost import TrustedHostMiddleware from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware import uuid from core.logging_config import setup_logging, request_id_var +from core.exception_handlers import register_exception_handlers from utils.files_whatcher import start_file_watcher from utils.pre_start_init import paths import threading @@ -24,22 +22,34 @@ from routes.ws_audio_transkrib import router as ws_audio_transkrib_router from routes.v1 import router as v1_router from routes.legacy import router as legacy_router -from models.fast_api_models import ErrorResponse import models from config import WS_DESCRIPTION logger = logging.getLogger(__name__) -class RequestIDMiddleware(BaseHTTPMiddleware): - async def dispatch(self, request: Request, call_next): +class RequestIDMiddleware: + def __init__(self, app): + self.app = app + + async def __call__(self, scope, receive, send): + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + request = Request(scope, receive) request_id = request.headers.get("X-Request-ID", str(uuid.uuid4())) request.state.request_id = request_id token = request_id_var.set(request_id) + + async def send_with_request_id(message): + if message["type"] == "http.response.start": + headers = MutableHeaders(raw=message["headers"]) + headers["X-Request-ID"] = request_id + await send(message) + try: - response = await call_next(request) - response.headers["X-Request-ID"] = request_id - return response + await self.app(scope, receive, send_with_request_id) finally: request_id_var.reset(token) @@ -112,43 +122,6 @@ async def lifespan(app): description=WS_DESCRIPTION ) -@app.exception_handler(RequestValidationError) -async def validation_exception_handler(request: Request, exc: RequestValidationError): - return JSONResponse( - status_code=422, - content=ErrorResponse( - success=False, - error_description="Validation error", - details=str(exc), - ).model_dump() - ) - - -@app.exception_handler(StarletteHTTPException) -async def http_exception_handler(request: Request, exc: StarletteHTTPException): - return JSONResponse( - status_code=exc.status_code, - content=ErrorResponse( - success=False, - error_description=exc.detail, - ).model_dump() - ) - - -@app.exception_handler(Exception) -async def general_exception_handler(request: Request, exc: Exception): - import traceback - logger.error(f"Unhandled exception: {exc}\n{traceback.format_exc()}") - return JSONResponse( - status_code=500, - content=ErrorResponse( - success=False, - error_description="Internal server error", - details=str(exc), - ).model_dump() - ) - - # RequestID middleware app.add_middleware(RequestIDMiddleware) @@ -176,9 +149,12 @@ async def general_exception_handler(request: Request, exc: Exception): # GZip middleware app.add_middleware( GZipMiddleware, - minimum_size=1000, + minimum_size=500 ) +# Exception handlers +register_exception_handlers(app) + # Static files app.mount("/static", StaticFiles(directory="static"), name="static") diff --git a/models/fast_api_models.py b/models/fast_api_models.py index bd40395..275cf22 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -93,22 +93,6 @@ class UserBase(BaseModel): is_active: bool = True -class Token(BaseModel): - """ - Модель токена доступа. - """ - access_token: str - token_type: str = "bearer" - - -class TokenPayload(BaseModel): - """ - Полезная нагрузка JWT-токена. - """ - sub: Optional[str] = None - exp: Optional[int] = None - - class SyncASRRequest(BaseModel): """ :parameter keep_raw: - Если False, то запрос вернёт только пост-обработанные данные do_punctuation и do_dialogue. diff --git a/requirements.txt b/requirements.txt index d59a477..e9f1ab4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -49,4 +49,8 @@ onnx-asr == 0.10.2 hf_xet onnxruntime -starlette \ No newline at end of file +starlette + +# For Auth / JWT +pyjwt~=2.10.1 +bcrypt~=4.3.0 diff --git a/tests/test_gzip_middleware.py b/tests/test_gzip_middleware.py index 3c2f872..0f7641d 100644 --- a/tests/test_gzip_middleware.py +++ b/tests/test_gzip_middleware.py @@ -6,6 +6,8 @@ @pytest.fixture def client(): + # Сбрасываем кэш middleware stack, чтобы перестроить с актуальными параметрами + app.middleware_stack = None with TestClient(app) as c: yield c @@ -26,5 +28,40 @@ def test_gzip_middleware_compresses_large_response(client): def test_gzip_middleware_skips_small_response(client): """Маленькие ответы (<1000 байт) не должны сжиматься.""" response = client.get("/", headers={"Accept-Encoding": "gzip"}) + print(f"[test_gzip_skips] request headers: {dict(response.request.headers)}") + print(f"[test_gzip_skips] status: {response.status_code}") + print(f"[test_gzip_skips] response headers: {dict(response.headers)}") + print(f"[test_gzip_skips] content length: {len(response.content)}") + print(f"[test_gzip_skips] text: {response.text[:500]}") assert response.status_code == 200 - assert "content-encoding" not in response.headers + print(f"[response.headers]: {response.headers}") + assert "content-encoding" not in response.headers, f"Unexpected content-encoding in headers: {dict(response.headers)}" + + +# def test_gzip_middleware_config_debug(): +# """Диагностика: выводим все параметры GZipMiddleware.""" +# gzip_middlewares = [m for m in app.user_middleware if m.cls is GZipMiddleware] +# print(f"[gzip debug] Найдено экземпляров GZipMiddleware: {len(gzip_middlewares)}") +# for i, m in enumerate(gzip_middlewares): +# opts = getattr(m, "kwargs", {}) +# print(f"[gzip debug] Экземпляр {i}: {opts}") +# assert len(gzip_middlewares) == 1 +# opts = getattr(gzip_middlewares[0], "kwargs", {}) +# assert opts.get("minimum_size") == 500 +# +# +# def test_gzip_skips_tiny_response(client): +# """Гарантированно маленький ответ не должен сжиматься.""" +# from main import app +# +# @app.get("/_test_tiny", include_in_schema=False) +# def _tiny(): +# return {"x": 1} +# +# resp = client.get("/_test_tiny", headers={"Accept-Encoding": "gzip"}) +# print(f"[tiny] status={resp.status_code}") +# print(f"[tiny] headers={dict(resp.headers)}") +# print(f"[tiny] len(content)={len(resp.content)}") +# print(f"[tiny] text={resp.text}") +# assert resp.status_code == 200 +# assert "content-encoding" not in resp.headers, f"Сжался ответ в {len(resp.content)} байт!" diff --git a/tests/test_security.py b/tests/test_security.py new file mode 100644 index 0000000..888d201 --- /dev/null +++ b/tests/test_security.py @@ -0,0 +1,105 @@ +import pytest +from datetime import timedelta + +from core.security import ( + get_password_hash, + verify_password, + create_access_token, + create_refresh_token, + decode_token, +) +from core.exceptions import ( + CredentialsException, + TokenExpiredException, + InvalidTokenException, +) + + +def test_get_password_hash_returns_string(): + """Хеш должен быть строкой и отличаться от исходного пароля.""" + password = "super_secret_123" + hashed = get_password_hash(password) + assert isinstance(hashed, str) + assert hashed != password + assert hashed.startswith("$2") + + +def test_verify_password_correct(): + """Верификация правильного пароля должна возвращать True.""" + password = "my_password" + hashed = get_password_hash(password) + assert verify_password(password, hashed) is True + + +def test_verify_password_incorrect(): + """Верификация неправильного пароля должна возвращать False.""" + password = "correct_horse_battery_staple" + hashed = get_password_hash(password) + assert verify_password("wrong_password", hashed) is False + + +def test_hash_salting_produces_different_hashes(): + """Один и тот же пароль должен давать разные хеши из-за соли.""" + password = "same_password" + hash1 = get_password_hash(password) + hash2 = get_password_hash(password) + assert hash1 != hash2 + assert verify_password(password, hash1) is True + assert verify_password(password, hash2) is True + + +def test_create_and_decode_access_token(): + """Access-токен создаётся и корректно декодируется.""" + token = create_access_token(data={"sub": "user123"}) + payload = decode_token(token, expected_type="access") + assert payload.sub == "user123" + assert payload.type == "access" + + +def test_create_and_decode_refresh_token(): + """Refresh-токен создаётся и корректно декодируется.""" + token = create_refresh_token(data={"sub": "user456"}) + payload = decode_token(token, expected_type="refresh") + assert payload.sub == "user456" + assert payload.type == "refresh" + + +def test_decode_token_wrong_type_raises(): + """Декодирование access-токена как refresh вызывает InvalidTokenException.""" + access_token = create_access_token(data={"sub": "user"}) + with pytest.raises(InvalidTokenException): + decode_token(access_token, expected_type="refresh") + + +def test_expired_token_raises(): + """Просроченный токен вызывает TokenExpiredException.""" + token = create_access_token(data={"sub": "user"}, expires_delta=timedelta(seconds=-1)) + with pytest.raises(TokenExpiredException): + decode_token(token) + + +def test_invalid_token_raises(): + """Совершенно невалидная строка токена вызывает InvalidTokenException.""" + with pytest.raises(InvalidTokenException): + decode_token("totally.invalid.token") + + +def test_credentials_exception_properties(): + """CredentialsException должен иметь статус 401 и корректный detail.""" + exc = CredentialsException() + assert exc.status_code == 401 + assert exc.detail == "Could not validate credentials" + + +def test_token_expired_exception_properties(): + """TokenExpiredException должен иметь статус 401 и корректный detail.""" + exc = TokenExpiredException() + assert exc.status_code == 401 + assert exc.detail == "Token has expired" + + +def test_invalid_token_exception_properties(): + """InvalidTokenException должен иметь статус 401 и корректный detail.""" + exc = InvalidTokenException() + assert exc.status_code == 401 + assert exc.detail == "Invalid token" From 45b3c91d086245f52e2c8cc0e21720a209ef70ef Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 4 May 2026 15:25:11 +0300 Subject: [PATCH 20/46] =?UTF-8?q?=D0=B0=D0=BA=D1=82=D1=83=D0=B0=D0=BB?= =?UTF-8?q?=D0=B8=D0=B7=D0=B0=D1=86=D0=B8=D1=8F=20=D1=82=D0=B5=D1=81=D1=82?= =?UTF-8?q?=D0=BE=D0=B2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/test_http_routes.py | 46 +++++++++++++++++++-------- tests/test_trusted_host_middleware.py | 10 ++++-- 2 files changed, 40 insertions(+), 16 deletions(-) diff --git a/tests/test_http_routes.py b/tests/test_http_routes.py index c43c0d9..120bbda 100644 --- a/tests/test_http_routes.py +++ b/tests/test_http_routes.py @@ -2,10 +2,12 @@ import json import httpx +import pytest +from config import settings from models.fast_api_models import V1BaseResponse as BaseResponse -BASE_URL = "http://127.0.0.1:49153/v1" +BASE_URL = f"http://127.0.0.1:{settings.PORT}/v1" async def _assert_base_response(body: dict, expect_success: bool): @@ -13,6 +15,7 @@ async def _assert_base_response(body: dict, expect_success: bool): assert body["success"] is expect_success +@pytest.mark.asyncio async def test_post_by_url_success(): async with httpx.AsyncClient() as client: payload = { @@ -22,9 +25,12 @@ async def test_post_by_url_success(): "do_dialogue": False, "do_punctuation": False, } - resp = await client.post( - f"{BASE_URL}/post_one_step_req", json=payload, timeout=120.0 - ) + try: + resp = await client.post( + f"{BASE_URL}/post_one_step_req", json=payload, timeout=120.0 + ) + except httpx.ConnectError: + pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") print("[POST /v1/post_one_step_req] status:", resp.status_code) body = resp.json() print(json.dumps(body, indent=2, ensure_ascii=False)) @@ -33,13 +39,17 @@ async def test_post_by_url_success(): return body +@pytest.mark.asyncio async def test_post_by_url_validation_error(): """Ожидаем ErrorResponse при 422.""" async with httpx.AsyncClient() as client: payload = {"keep_raw": True} # отсутствует AudioFileUrl - resp = await client.post( - f"{BASE_URL}/post_one_step_req", json=payload, timeout=10.0 - ) + try: + resp = await client.post( + f"{BASE_URL}/post_one_step_req", json=payload, timeout=10.0 + ) + except httpx.ConnectError: + pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") print("[POST /v1/post_one_step_req 422] status:", resp.status_code) body = resp.json() print(json.dumps(body, indent=2, ensure_ascii=False)) @@ -50,9 +60,10 @@ async def test_post_by_url_validation_error(): return body +@pytest.mark.asyncio async def test_post_by_file_success(): async with httpx.AsyncClient() as client: - with open("../examples/orig.wav", "rb") as f: + with open("./examples/orig.wav", "rb") as f: files = {"file": ("orig.wav", f, "audio/wav")} data = { "keep_raw": "true", @@ -62,9 +73,12 @@ async def test_post_by_file_success(): "do_diarization": "false", "diar_vad_sensity": "3", } - resp = await client.post( - f"{BASE_URL}/post_file", data=data, files=files, timeout=120.0 - ) + try: + resp = await client.post( + f"{BASE_URL}/post_file", data=data, files=files, timeout=20.0 + ) + except httpx.ConnectError: + pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") print("[POST /v1/post_file] status:", resp.status_code) body = resp.json() print(json.dumps(body, indent=2, ensure_ascii=False)) @@ -73,13 +87,17 @@ async def test_post_by_file_success(): return body +@pytest.mark.asyncio async def test_post_by_file_validation_error(): """Ожидаем ErrorResponse при 422 (нет файла).""" async with httpx.AsyncClient() as client: data = {"keep_raw": "true"} - resp = await client.post( - f"{BASE_URL}/post_file", data=data, timeout=10.0 - ) + try: + resp = await client.post( + f"{BASE_URL}/post_file", data=data, timeout=10.0 + ) + except httpx.ConnectError: + pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") print("[POST /v1/post_file 422] status:", resp.status_code) body = resp.json() print(json.dumps(body, indent=2, ensure_ascii=False)) diff --git a/tests/test_trusted_host_middleware.py b/tests/test_trusted_host_middleware.py index 522a05d..fdf61ac 100644 --- a/tests/test_trusted_host_middleware.py +++ b/tests/test_trusted_host_middleware.py @@ -10,20 +10,26 @@ def restricted_client(): original_hosts = settings.ALLOWED_HOSTS.copy() settings.ALLOWED_HOSTS[:] = ["trusted.example.com"] # Сбрасываем кэш middleware stack, чтобы перестроить с актуальными настройками - app._middleware_stack = None + app.middleware_stack = None with TestClient(app) as client: yield client settings.ALLOWED_HOSTS[:] = original_hosts - app._middleware_stack = None + app.middleware_stack = None def test_trusted_host_middleware_rejects_invalid_host(restricted_client): """Запросы с недоверенным Host должны отклоняться с 400.""" + print(f"[trusted_host_reject] allowed_hosts={settings.ALLOWED_HOSTS}") response = restricted_client.get("/", headers={"Host": "untrusted.example.com"}) + print(f"[trusted_host_reject] request headers: {dict(response.request.headers)}") + print(f"[trusted_host_reject] response status: {response.status_code}") + print(f"[trusted_host_reject] response headers: {dict(response.headers)}") assert response.status_code == 400 def test_trusted_host_middleware_allows_valid_host(restricted_client): """Запросы с доверенным Host должны проходить (не 400).""" + print(f"[trusted_host_allow] allowed_hosts={settings.ALLOWED_HOSTS}") response = restricted_client.get("/", headers={"Host": "trusted.example.com"}) + print(f"[trusted_host_allow] response status: {response.status_code}") assert response.status_code != 400 From 31e7725504163f6f59b0b189630c895c019b439c Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 4 May 2026 17:12:27 +0300 Subject: [PATCH 21/46] =?UTF-8?q?=D0=92=D0=B2=D0=B5=D0=B4=D0=B5=D0=BD?= =?UTF-8?q?=D0=B8=D0=B5=20=D0=BC=D0=BE=D0=B4=D0=B5=D0=BB=D0=B5=D0=B9=20?= =?UTF-8?q?=D0=BF=D0=BE=D0=BB=D1=8C=D0=B7=D0=BE=D0=B2=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=BB=D0=B5=D0=B9=20=D0=B8=20=D0=BF=D0=BB=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=B6=D0=B5=D0=B9.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- models/domain/__init__.py | 15 +++++++ models/domain/audit.py | 55 +++++++++++++++++++++++ models/domain/billing.py | 94 +++++++++++++++++++++++++++++++++++++++ models/domain/user.py | 85 +++++++++++++++++++++++++++++++++++ models/enums.py | 37 +++++++++++++++ 5 files changed, 286 insertions(+) create mode 100644 models/domain/__init__.py create mode 100644 models/domain/audit.py create mode 100644 models/domain/billing.py create mode 100644 models/domain/user.py create mode 100644 models/enums.py diff --git a/models/domain/__init__.py b/models/domain/__init__.py new file mode 100644 index 0000000..75a8c3e --- /dev/null +++ b/models/domain/__init__.py @@ -0,0 +1,15 @@ +from models.domain.user import User, UserAccount, QuotaInfo, UserProfileResponse +from models.domain.billing import Subscription, Transaction, RecurringPayment +from models.domain.audit import ASRUsageLog, AdminActionLog + +__all__ = [ + "User", + "UserAccount", + "QuotaInfo", + "UserProfileResponse", + "Subscription", + "Transaction", + "RecurringPayment", + "ASRUsageLog", + "AdminActionLog", +] diff --git a/models/domain/audit.py b/models/domain/audit.py new file mode 100644 index 0000000..adb3e6c --- /dev/null +++ b/models/domain/audit.py @@ -0,0 +1,55 @@ +from datetime import datetime +from decimal import Decimal +from typing import Optional + +from pydantic import BaseModel, Field + + +class ASRUsageLog(BaseModel): + """Лог использования ASR для учёта квот и биллинга.""" + id: str + user_id: str + endpoint: str = Field(description="Вызываемый эндпоинт") + audio_duration_sec: Optional[float] = Field(default=None, ge=0) + request_size_bytes: int = Field(default=0, ge=0) + timestamp: Optional[datetime] = None + cost: Optional[Decimal] = Field(default=None, decimal_places=2) + quota_consumed: int = Field(default=0, ge=0) + + model_config = { + "json_schema_extra": { + "example": { + "id": "log_12345", + "user_id": "usr_12345", + "endpoint": "/api/v1/asr/url", + "audio_duration_sec": 120.5, + "request_size_bytes": 1048576, + "timestamp": "2024-01-01T12:00:00Z", + "cost": "15.50", + "quota_consumed": 1, + } + } + } + + +class AdminActionLog(BaseModel): + """Лог административных действий для аудита.""" + id: str + admin_id: str + action: str = Field(description="Выполненное действие") + target_user_id: Optional[str] = Field(default=None, description="ID целевого пользователя") + details: dict = Field(default_factory=dict, description="Дополнительные параметры") + timestamp: Optional[datetime] = None + + model_config = { + "json_schema_extra": { + "example": { + "id": "adm_12345", + "admin_id": "usr_admin", + "action": "block_user", + "target_user_id": "usr_12345", + "details": {"reason": "violation"}, + "timestamp": "2024-01-01T12:00:00Z", + } + } + } diff --git a/models/domain/billing.py b/models/domain/billing.py new file mode 100644 index 0000000..80cac80 --- /dev/null +++ b/models/domain/billing.py @@ -0,0 +1,94 @@ +from datetime import datetime +from decimal import Decimal +from typing import Optional + +from pydantic import BaseModel, Field + +from models.enums import SubscriptionType, SubscriptionStatus, PaymentMethod, TransactionStatus + + +class Subscription(BaseModel): + """Модель подписки пользователя.""" + id: str + user_id: str + type: SubscriptionType + start_date: Optional[datetime] = None + end_date: Optional[datetime] = None + status: SubscriptionStatus = SubscriptionStatus.active + payment_method: Optional[PaymentMethod] = None + + model_config = { + "json_schema_extra": { + "example": { + "id": "sub_12345", + "user_id": "usr_12345", + "type": "pro", + "start_date": "2024-01-01T00:00:00Z", + "end_date": "2024-02-01T00:00:00Z", + "status": "active", + "payment_method": "card", + } + } + } + + +class Transaction(BaseModel): + """Модель разового платежа.""" + id: str + user_id: str + amount: Decimal = Field(decimal_places=2, gt=0) + currency: str = Field(default="RUB", max_length=3) + status: TransactionStatus + external_id: Optional[str] = Field(default=None, description="ID платежа во внешней системе") + payment_method: PaymentMethod + is_recurring: bool = False + created_at: Optional[datetime] = None + + model_config = { + "json_schema_extra": { + "example": { + "id": "txn_12345", + "user_id": "usr_12345", + "amount": "999.00", + "currency": "RUB", + "status": "completed", + "external_id": "ext_12345", + "payment_method": "card", + "is_recurring": False, + "created_at": "2024-01-01T12:00:00Z", + } + } + } + + +class RecurringPayment(BaseModel): + """Модель рекуррентного (автоматического) платежа.""" + id: str + user_id: str + subscription_id: str + amount: Decimal = Field(decimal_places=2, gt=0) + currency: str = Field(default="RUB", max_length=3) + interval_days: int = Field(default=30, ge=1, description="Периодичность списания в днях") + next_payment_date: Optional[datetime] = None + status: SubscriptionStatus = SubscriptionStatus.active + payment_method: PaymentMethod + external_subscription_id: Optional[str] = Field( + default=None, description="ID подписки в платёжном шлюзе" + ) + + model_config = { + "json_schema_extra": { + "example": { + "id": "rec_12345", + "user_id": "usr_12345", + "subscription_id": "sub_12345", + "amount": "999.00", + "currency": "RUB", + "interval_days": 30, + "next_payment_date": "2024-02-01T00:00:00Z", + "status": "active", + "payment_method": "card", + "external_subscription_id": "ext_sub_12345", + } + } + } diff --git a/models/domain/user.py b/models/domain/user.py new file mode 100644 index 0000000..8852f8a --- /dev/null +++ b/models/domain/user.py @@ -0,0 +1,85 @@ +from datetime import datetime +from decimal import Decimal +from typing import Optional + +from pydantic import BaseModel, Field + +from models.enums import Role, SubscriptionType +from models.domain.billing import Subscription + + +class User(BaseModel): + """Доменная модель пользователя.""" + id: str + email: Optional[str] = None + hashed_password: Optional[str] = None + role: Role = Role.user + is_active: bool = True + daily_quota: int = Field(default=10, ge=0, description="Дневная квота запросов") + quota_used_today: int = Field(default=0, ge=0, description="Использовано сегодня") + subscription_type: Optional[SubscriptionType] = None + subscription_expires: Optional[datetime] = None + created_at: Optional[datetime] = None + updated_at: Optional[datetime] = None + + model_config = { + "json_schema_extra": { + "example": { + "id": "usr_12345", + "email": "user@example.com", + "role": "user", + "is_active": True, + "daily_quota": 10, + "quota_used_today": 0, + "subscription_type": None, + "subscription_expires": None, + } + } + } + + +class UserAccount(BaseModel): + """Модель аккаунта пользователя с балансом и подпиской.""" + user_id: str + balance: Decimal = Field(default=Decimal("0.00"), decimal_places=2) + currency: str = Field(default="RUB", max_length=3) + tariff: str = Field(default="free") + subscription: Optional[Subscription] = None + auto_renew: bool = False + + model_config = { + "json_schema_extra": { + "example": { + "user_id": "usr_12345", + "balance": "0.00", + "currency": "RUB", + "tariff": "free", + "subscription": None, + "auto_renew": False, + } + } + } + + +class QuotaInfo(BaseModel): + """Информация о текущей квоте пользователя.""" + used: int = Field(ge=0, description="Использовано запросов") + limit: int = Field(ge=0, description="Лимит запросов") + reset_at: datetime = Field(description="Время сброса квоты") + + model_config = { + "json_schema_extra": { + "example": { + "used": 5, + "limit": 10, + "reset_at": "2024-01-02T00:00:00Z", + } + } + } + + +class UserProfileResponse(BaseModel): + """Унифицированный ответ с профилем пользователя.""" + user: User + account: UserAccount + quota: QuotaInfo diff --git a/models/enums.py b/models/enums.py new file mode 100644 index 0000000..d1ff0f0 --- /dev/null +++ b/models/enums.py @@ -0,0 +1,37 @@ +from enum import Enum + + +class Role(str, Enum): + """Роли пользователей для RBAC.""" + guest = "guest" + user = "user" + admin = "admin" + superadmin = "superadmin" + + +class SubscriptionType(str, Enum): + """Типы подписок.""" + free = "free" + pro = "pro" + enterprise = "enterprise" + + +class SubscriptionStatus(str, Enum): + """Статусы подписки.""" + active = "active" + expired = "expired" + cancelled = "cancelled" + + +class TransactionStatus(str, Enum): + """Статусы транзакции.""" + pending = "pending" + completed = "completed" + failed = "failed" + + +class PaymentMethod(str, Enum): + """Способы оплаты.""" + card = "card" + crypto = "crypto" + bank_transfer = "bank_transfer" From 26517edbbf8acd671aa2c7dbca9ee2c8bacfde1a Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 4 May 2026 17:22:18 +0300 Subject: [PATCH 22/46] =?UTF-8?q?=D0=92=D0=B2=D0=B5=D0=B4=D0=B5=D0=BD?= =?UTF-8?q?=D0=B8=D0=B5=20=D0=BC=D0=BE=D0=B4=D0=B5=D0=BB=D0=B5=D0=B9=20?= =?UTF-8?q?=D0=BF=D0=BE=D0=BB=D1=8C=D0=B7=D0=BE=D0=B2=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=BB=D0=B5=D0=B9=20=D0=B8=20=D0=BF=D0=BB=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=B6=D0=B5=D0=B9.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.py | 3 +++ models/domain/user.py | 7 ++++++- 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/config.py b/config.py index cfdcf59..f2d50ff 100644 --- a/config.py +++ b/config.py @@ -111,6 +111,9 @@ class Settings(BaseSettings): REFRESH_TOKEN_EXPIRE_MINUTES: int = 10080 # 7 дней BCRYPT_ROUNDS: int = 12 + # Quota settings + GUEST_DAILY_QUOTA: int = 10 + @field_validator( 'IS_PROD', 'MAKE_MONO', 'USE_BATCH', 'VAD_WITH_GPU', 'CAN_PUNCTUATE', 'PUNCTUATE_WITH_GPU', 'CAN_DIAR', diff --git a/models/domain/user.py b/models/domain/user.py index 8852f8a..d4d15f4 100644 --- a/models/domain/user.py +++ b/models/domain/user.py @@ -4,6 +4,7 @@ from pydantic import BaseModel, Field +from config import settings from models.enums import Role, SubscriptionType from models.domain.billing import Subscription @@ -15,7 +16,11 @@ class User(BaseModel): hashed_password: Optional[str] = None role: Role = Role.user is_active: bool = True - daily_quota: int = Field(default=10, ge=0, description="Дневная квота запросов") + daily_quota: int = Field( + default=settings.GUEST_DAILY_QUOTA, + ge=0, + description="Дневная квота запросов" + ) quota_used_today: int = Field(default=0, ge=0, description="Использовано сегодня") subscription_type: Optional[SubscriptionType] = None subscription_expires: Optional[datetime] = None From 7239ff04e502ec956dcb328f96fec9cb70c1bea5 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 4 May 2026 17:35:40 +0300 Subject: [PATCH 23/46] =?UTF-8?q?=D0=92=D0=B2=D0=B5=D0=B4=D0=B5=D0=BD?= =?UTF-8?q?=D0=B8=D0=B5=20=D0=BC=D0=BE=D0=B4=D0=B5=D0=BB=D0=B5=D0=B9=20?= =?UTF-8?q?=D0=BF=D0=BE=D0=BB=D1=8C=D0=B7=D0=BE=D0=B2=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=BB=D0=B5=D0=B9=20=D0=B8=20=D0=BF=D0=BB=D0=B0=D1=82=D0=B5?= =?UTF-8?q?=D0=B6=D0=B5=D0=B9.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/deps.py | 99 ++++++++++++++++++++++++++++++++++++++++++++++ core/auth.py | 24 +++++++++++ core/exceptions.py | 12 ++++++ 3 files changed, 135 insertions(+) create mode 100644 api/deps.py create mode 100644 core/auth.py diff --git a/api/deps.py b/api/deps.py new file mode 100644 index 0000000..c4f1c60 --- /dev/null +++ b/api/deps.py @@ -0,0 +1,99 @@ +from datetime import datetime, timezone +from typing import Optional + +from fastapi import Depends +from fastapi.security import OAuth2PasswordBearer + +from config import settings +from core.exceptions import ( + CredentialsException, + PermissionDeniedException, + RateLimitExceededException, +) +from core.security import decode_token +from models.domain.user import User +from models.enums import Role, SubscriptionType + + +oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/v1/auth/login", auto_error=False) + + +async def get_current_user(token: Optional[str] = Depends(oauth2_scheme)) -> User: + """Получение текущего пользователя из токена или гостевого доступа.""" + if not token: + return User( + id="guest", + role=Role.guest, + daily_quota=settings.GUEST_DAILY_QUOTA, + quota_used_today=0, + is_active=True, + ) + + try: + payload = decode_token(token, expected_type="access") + except Exception: + raise CredentialsException() + + # Заглушка: в реальности роль и квота берутся из БД + return User( + id=payload.sub, + role=Role.user, + daily_quota=settings.GUEST_DAILY_QUOTA, + quota_used_today=0, + is_active=True, + ) + + +async def get_current_active_user( + current_user: User = Depends(get_current_user), +) -> User: + """Проверка, что пользователь активен.""" + if not current_user.is_active: + raise CredentialsException() + return current_user + + +def require_role(*roles: Role): + """Зависимость для проверки роли пользователя.""" + async def role_checker( + current_user: User = Depends(get_current_active_user), + ) -> User: + if current_user.role not in roles: + raise PermissionDeniedException() + return current_user + + return Depends(role_checker) + + +async def check_daily_quota( + current_user: User = Depends(get_current_user), +) -> User: + """Проверка дневной квоты пользователя.""" + # Админы и суперадмины не ограничены + if current_user.role in (Role.admin, Role.superadmin): + return current_user + + # Пользователи с активной подпиской pro/enterprise не ограничены + if current_user.subscription_type in ( + SubscriptionType.pro, + SubscriptionType.enterprise, + ): + if ( + current_user.subscription_expires is None + or current_user.subscription_expires > datetime.now(timezone.utc) + ): + return current_user + + # Обычные пользователи и гости проверяются по квоте + if current_user.quota_used_today >= current_user.daily_quota: + raise RateLimitExceededException() + + return current_user + + +async def require_paid_access( + current_user: User = Depends(get_current_user), +) -> User: + """Заглушка для проверки разовых/рекуррентных платежей.""" + # TODO: реализовать проверку оплаты при подключении платёжной системы + return current_user diff --git a/core/auth.py b/core/auth.py new file mode 100644 index 0000000..de89cc3 --- /dev/null +++ b/core/auth.py @@ -0,0 +1,24 @@ +from typing import Optional + +from models.domain.user import User +from models.enums import Role +from core.security import get_password_hash + + +async def authenticate_user(email: str, password: str) -> Optional[User]: + """Аутентификация пользователя по email и паролю (заглушка).""" + # TODO: реализовать проверку через БД + return None + + +async def register_user(email: str, password: str) -> User: + """Регистрация нового пользователя (заглушка).""" + # TODO: реализовать сохранение в БД + hashed_password = get_password_hash(password) + return User( + id="usr_new_12345", + email=email, + hashed_password=hashed_password, + role=Role.user, + is_active=True, + ) diff --git a/core/exceptions.py b/core/exceptions.py index dfdbb60..2e6aa1e 100644 --- a/core/exceptions.py +++ b/core/exceptions.py @@ -17,3 +17,15 @@ class InvalidTokenException(HTTPException): """401 — невалидный токен.""" def __init__(self): super().__init__(status_code=401, detail="Invalid token") + + +class PermissionDeniedException(HTTPException): + """403 — доступ запрещён.""" + def __init__(self): + super().__init__(status_code=403, detail="Permission denied") + + +class RateLimitExceededException(HTTPException): + """429 — превышена дневная квота.""" + def __init__(self): + super().__init__(status_code=429, detail="Daily quota exceeded") From 0a7dd1f8805f2b2e66caa6c182c95823661526c5 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Tue, 5 May 2026 12:46:12 +0300 Subject: [PATCH 24/46] =?UTF-8?q?=D0=98=D0=B7=D0=BC=D0=B5=D0=BD=D0=B5?= =?UTF-8?q?=D0=BD=D0=B8=D0=B5=20=D1=84=D0=B0=D0=B9=D0=BB=D0=BE=D0=B2=D0=BE?= =?UTF-8?q?=D0=B9=20=D1=81=D1=82=D1=80=D1=83=D0=BA=D1=82=D1=83=D1=80=D1=8B?= =?UTF-8?q?=20API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/engine/file_recognition.py | 7 ++- Recognizer/engine/stream_recognition.py | 5 +- {routes => api}/legacy/__init__.py | 0 {routes => api}/legacy/demo_page.py | 0 {routes => api}/legacy/is_alive.py | 0 {routes => api}/legacy/post_by_file_FORM.py | 0 {routes => api}/legacy/post_by_url.py | 0 {routes => api}/legacy/root.py | 1 - api/v1/__init__.py | 1 + api/v1/api.py | 15 +++++ api/v1/endpoints/__init__.py | 13 ++++ .../v1/endpoints/asr_file.py | 40 ++++++------- .../v1/endpoints/asr_url.py | 47 ++++++++------- .../v1/endpoints/asr_ws.py | 48 +++++---------- api/v1/endpoints/health.py | 59 +++++++++++++++++++ api/v1/endpoints/root.py | 31 ++++++++++ main.py | 6 +- models/fast_api_models.py | 10 ++-- routes/__init__.py | 1 - routes/v1/__init__.py | 11 ---- routes/v1/is_alive.py | 35 ----------- routes/v1/root.py | 28 --------- templates/index.html | 8 +-- tests/test_http_routes.py | 10 ++-- 24 files changed, 205 insertions(+), 171 deletions(-) rename {routes => api}/legacy/__init__.py (100%) rename {routes => api}/legacy/demo_page.py (100%) rename {routes => api}/legacy/is_alive.py (100%) rename {routes => api}/legacy/post_by_file_FORM.py (100%) rename {routes => api}/legacy/post_by_url.py (100%) rename {routes => api}/legacy/root.py (99%) create mode 100644 api/v1/__init__.py create mode 100644 api/v1/api.py create mode 100644 api/v1/endpoints/__init__.py rename routes/v1/post_by_file_FORM.py => api/v1/endpoints/asr_file.py (79%) rename routes/v1/post_by_url.py => api/v1/endpoints/asr_url.py (55%) rename routes/ws_audio_transkrib.py => api/v1/endpoints/asr_ws.py (83%) create mode 100644 api/v1/endpoints/health.py create mode 100644 api/v1/endpoints/root.py delete mode 100644 routes/__init__.py delete mode 100644 routes/v1/__init__.py delete mode 100644 routes/v1/is_alive.py delete mode 100644 routes/v1/root.py diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index 2b15457..ce3cc7f 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -141,9 +141,12 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): # Снижаем скорость аудио по необходимости if params.do_auto_speech_speed_correction or params.speech_speed_correction_multiplier != 1: logger.debug("Будут использованы механизмы анализа скорости речи и замедления аудио") - asr_result_wo_conf, speed, multiplier = asyncio.run(recognise_w_speed_correction(audio_asr, + asr_result_wo_conf, speed, multiplier = asyncio.run(recognise_w_speed_correction( + audio_data=audio_asr, can_slow_down=True, - multiplier=params.speech_speed_correction_multiplier)) + multiplier=params.speech_speed_correction_multiplier, + recognizer=recognizer) + ) params.speech_speed_correction_multiplier = multiplier else: # Производим распознавание diff --git a/Recognizer/engine/stream_recognition.py b/Recognizer/engine/stream_recognition.py index fda8780..0e033ab 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -54,10 +54,11 @@ async def simple_recognise(audio_data, recognizer) -> dict: async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_slow_down = False, - ) -> tuple: + recognizer = None) -> tuple: """ Распознавание чанка с возможностью контроля быстрой речи. + :param recognizer: Обязательно получить класс модели распознавания. :param multiplier: Float :param can_slow_down: Boolean :param audio_data: Аудиоданные в формате Audiosegment (puDub). @@ -71,7 +72,7 @@ async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_sl # Парсим результат - result = await simple_recognise(audio_data) + result = await simple_recognise(audio_data, recognizer) if can_slow_down and multiplier == 1: speed = calc_speed(result) diff --git a/routes/legacy/__init__.py b/api/legacy/__init__.py similarity index 100% rename from routes/legacy/__init__.py rename to api/legacy/__init__.py diff --git a/routes/legacy/demo_page.py b/api/legacy/demo_page.py similarity index 100% rename from routes/legacy/demo_page.py rename to api/legacy/demo_page.py diff --git a/routes/legacy/is_alive.py b/api/legacy/is_alive.py similarity index 100% rename from routes/legacy/is_alive.py rename to api/legacy/is_alive.py diff --git a/routes/legacy/post_by_file_FORM.py b/api/legacy/post_by_file_FORM.py similarity index 100% rename from routes/legacy/post_by_file_FORM.py rename to api/legacy/post_by_file_FORM.py diff --git a/routes/legacy/post_by_url.py b/api/legacy/post_by_url.py similarity index 100% rename from routes/legacy/post_by_url.py rename to api/legacy/post_by_url.py diff --git a/routes/legacy/root.py b/api/legacy/root.py similarity index 99% rename from routes/legacy/root.py rename to api/legacy/root.py index a12b080..084ae1c 100644 --- a/routes/legacy/root.py +++ b/api/legacy/root.py @@ -22,4 +22,3 @@ async def root(): "/demo": "DEMO UI page" }, "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs"} - diff --git a/api/v1/__init__.py b/api/v1/__init__.py new file mode 100644 index 0000000..17d7665 --- /dev/null +++ b/api/v1/__init__.py @@ -0,0 +1 @@ +# api/v1 package diff --git a/api/v1/api.py b/api/v1/api.py new file mode 100644 index 0000000..0d864ed --- /dev/null +++ b/api/v1/api.py @@ -0,0 +1,15 @@ +from fastapi import APIRouter + +from api.v1.endpoints.root import router as root_router +from api.v1.endpoints.asr_url import router as asr_url_router +from api.v1.endpoints.asr_file import router as asr_file_router +from api.v1.endpoints.asr_ws import router as asr_ws_router +from api.v1.endpoints.health import router as health_router + +router = APIRouter(prefix="/api/v1") + +router.include_router(root_router) +router.include_router(asr_url_router) +router.include_router(asr_file_router) +router.include_router(asr_ws_router) +router.include_router(health_router) diff --git a/api/v1/endpoints/__init__.py b/api/v1/endpoints/__init__.py new file mode 100644 index 0000000..4411708 --- /dev/null +++ b/api/v1/endpoints/__init__.py @@ -0,0 +1,13 @@ +from api.v1.endpoints.root import router as root_router +from api.v1.endpoints.asr_url import router as asr_url_router +from api.v1.endpoints.asr_file import router as asr_file_router +from api.v1.endpoints.asr_ws import router as asr_ws_router +from api.v1.endpoints.health import router as health_router + +__all__ = [ + "root_router", + "asr_url_router", + "asr_file_router", + "asr_ws_router", + "health_router", +] diff --git a/routes/v1/post_by_file_FORM.py b/api/v1/endpoints/asr_file.py similarity index 79% rename from routes/v1/post_by_file_FORM.py rename to api/v1/endpoints/asr_file.py index 413b75d..990fd58 100644 --- a/routes/v1/post_by_file_FORM.py +++ b/api/v1/endpoints/asr_file.py @@ -1,8 +1,9 @@ from io import BytesIO import asyncio +import logging from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile -from models.fast_api_models import PostFileRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData +from models.fast_api_models import PostFileRequest, V1BaseResponse, ASRData, RawData, SentencedData, DiarizedData from Recognizer.engine.file_recognition import process_file from Recognizer import get_recognizer, Recognizer from Punctuation import get_punctuator, SbertPuncCaseOnnx @@ -10,10 +11,9 @@ from Diarisation import get_diarizer from Diarisation.do_diarize import Diarizer -import logging - logger = logging.getLogger(__name__) -router = APIRouter() +router = APIRouter(prefix="/asr", tags=["ASR"]) + # Функция для извлечения параметров из FormData def get_file_request( @@ -39,28 +39,26 @@ def get_file_request( diar_vad_sensity=diar_vad_sensity, batch_size=batch_size, make_mono=make_mono, - do_auto_speech_speed_correction = do_auto_speech_speed_correction, - speech_speed_correction_multiplier = speech_speed_correction_multiplier + do_auto_speech_speed_correction=do_auto_speech_speed_correction, + speech_speed_correction_multiplier=speech_speed_correction_multiplier ) -@router.post("/post_file", response_model=V1ASRResponse) +@router.post("/file", response_model=V1BaseResponse) async def async_receive_file( file: UploadFile = File(description="Аудиофайл для обработки"), params: PostFileRequest = Depends(get_file_request), recognizer: Recognizer = Depends(get_recognizer), punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), diarizer: Diarizer = Depends(get_diarizer) -) -> V1ASRResponse: - - +) -> V1BaseResponse: try: buffer = BytesIO(await file.read()) buffer.seek(0) except Exception as e: error_description = f"Не удалось сохранить файл для распознавания: {file.filename}, размер файла: {file.size}, по причине: {e}" logger.error(error_description) - return V1ASRResponse( + return V1BaseResponse( success=False, error_description=error_description, data=ASRData() @@ -68,17 +66,19 @@ async def async_receive_file( else: logger.info(f"Получен и сохранён файл {file.filename}") try: - result_dict = await asyncio.to_thread(process_file, - tmp_path=buffer, - params=params, - recognizer=recognizer, - punctuator=punctuator, - diarizer=diarizer) - return V1ASRResponse( + result_dict = await asyncio.to_thread( + process_file, + tmp_path=buffer, + params=params, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer + ) + return V1BaseResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), data=ASRData( - raw_data=RawData(**result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, + raw_data=RawData.model_validate(result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None ) @@ -86,7 +86,7 @@ async def async_receive_file( except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - return V1ASRResponse( + return V1BaseResponse( success=False, error_description=str(error_description), data=ASRData() diff --git a/routes/v1/post_by_url.py b/api/v1/endpoints/asr_url.py similarity index 55% rename from routes/v1/post_by_url.py rename to api/v1/endpoints/asr_url.py index c89b864..490a476 100644 --- a/routes/v1/post_by_url.py +++ b/api/v1/endpoints/asr_url.py @@ -1,8 +1,9 @@ import uuid import asyncio +import logging from utils.pre_start_init import posted_and_downloaded_audio from utils.get_audio_file import getting_audiofile, open_default_audiofile -from models.fast_api_models import SyncASRRequest, V1ASRResponse, ASRData, RawData, SentencedData, DiarizedData +from models.fast_api_models import SyncASRRequest, V1BaseResponse, ASRData, RawData, SentencedData, DiarizedData from fastapi import APIRouter, Depends from Recognizer import get_recognizer, Recognizer @@ -11,18 +12,18 @@ from Diarisation import get_diarizer from Diarisation.do_diarize import Diarizer - -import logging logger = logging.getLogger(__name__) -router = APIRouter() +router = APIRouter(prefix="/asr", tags=["ASR"]) -@router.post("/post_one_step_req", response_model=V1ASRResponse) -async def post_v1(params: SyncASRRequest, - recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), - diarizer: Diarizer = Depends(get_diarizer) - ) -> V1ASRResponse: + +@router.post("/url", response_model=V1BaseResponse) +async def post_v1( + params: SyncASRRequest, + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) +) -> V1BaseResponse: post_id = uuid.uuid4() if params.AudioFileUrl: res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) @@ -30,25 +31,29 @@ async def post_v1(params: SyncASRRequest, res, error_description = await open_default_audiofile(post_id) if not res: - logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') - return V1ASRResponse( + logger.error( + f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}' + ) + return V1BaseResponse( success=False, error_description=error_description, data=ASRData() ) try: - result_dict = await asyncio.to_thread(process_file, - tmp_path=posted_and_downloaded_audio[post_id], - params=params, - recognizer=recognizer, - punctuator=punctuator, - diarizer=diarizer) - return V1ASRResponse( + result_dict = await asyncio.to_thread( + process_file, + tmp_path=posted_and_downloaded_audio[post_id], + params=params, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer + ) + return V1BaseResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), data=ASRData( - raw_data=RawData(**result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, + raw_data=RawData.model_validate(result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None ) @@ -56,7 +61,7 @@ async def post_v1(params: SyncASRRequest, except Exception as e: error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) - return V1ASRResponse( + return V1BaseResponse( success=False, error_description=str(error_description), data=ASRData() diff --git a/routes/ws_audio_transkrib.py b/api/v1/endpoints/asr_ws.py similarity index 83% rename from routes/ws_audio_transkrib.py rename to api/v1/endpoints/asr_ws.py index bece710..c64b765 100644 --- a/routes/ws_audio_transkrib.py +++ b/api/v1/endpoints/asr_ws.py @@ -6,30 +6,29 @@ import uuid from io import BytesIO -from fastapi import APIRouter, WebSocket +from fastapi import APIRouter, WebSocket, Depends from utils.chunk_doing import find_last_speech_position -from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration,ws_collected_asr_res +from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration, ws_collected_asr_res from utils.send_messages import send_messages from utils.tokens_to_Result import process_single_token_vocab_output from utils.resamppling import async_resample_audiosegment -from fastapi import Depends from Recognizer import get_recognizer, Recognizer - from Recognizer.engine.sentensizer import do_sensitizing from Recognizer.engine.stream_recognition import simple_recognise from Punctuation import get_punctuator, SbertPuncCaseOnnx -router = APIRouter() +router = APIRouter(prefix="/asr", tags=["ASR"]) logger = logging.getLogger(__name__) @router.websocket("/ws") -async def websocket(ws: WebSocket, - recognizer: Recognizer = Depends(get_recognizer), - punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), - ): - wait_null_answers=True +async def websocket( + ws: WebSocket, + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), +): + wait_null_answers = True client_id = uuid.uuid4() logger.debug(f'Принят новый сокет id = {client_id}') audio_buffer[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) @@ -40,7 +39,7 @@ async def websocket(ws: WebSocket, do_dialogue = False do_punctuation = False audio_format = 'raw' - sample_rate = settings.BASE_SAMPLE_RATE # Если не получен фреймрейт в конфиге сокета, по попытается принять с конфигом модели. + sample_rate = settings.BASE_SAMPLE_RATE sentenced_data = None error_description = None @@ -81,27 +80,23 @@ async def websocket(ws: WebSocket, logger.error(f'Error text message compiling. Message:{message} - error:{e} in channel {channel_name}') elif isinstance(message, dict) and message.get('bytes'): try: - # Получаем новый чанк с данными chunk = message.get('bytes') if audio_format == 'pcm16': - # Проверяем и добавляем недостающие нулевые байты в чанки. if len(chunk) % 2 != 0: chunk += bytes(2 - (len(chunk) % 2)) - # Переводим чанк в объект Audiosegment audiosegment_chunk = AudioSegment( chunk, - frame_rate = sample_rate, # Частота дискретизации - sample_width = 2, # Ширина сэмпла (2 байта для int16) - channels = 1 # Количество каналов. По умолчанию - 1, Моно. + frame_rate=sample_rate, + sample_width=2, + channels=1 ) else: try: buffer = BytesIO(chunk) buffer.seek(0) - # buffer.write() audiosegment_chunk = AudioSegment.from_file(buffer) except Exception as e: @@ -109,18 +104,14 @@ async def websocket(ws: WebSocket, else: logger.debug(f"Чанк принят и распознан in channel {channel_name}") - # Приводим фреймрейт к фреймрейту модели if audiosegment_chunk.frame_rate != settings.BASE_SAMPLE_RATE: audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, settings.BASE_SAMPLE_RATE) if audiosegment_chunk.channels != 1: audiosegment_chunk = audiosegment_chunk.set_channels(1) - # Копим буфер audio_buffer[client_id] += audiosegment_chunk - # Накопили больше нормы - if (audio_overlap[client_id]+audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: - # Проверяем новый чанк перед объединением (там же режем хвост и добавляем его при необходимости) + if (audio_overlap[client_id] + audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: await find_last_speech_position(client_id, is_last_chunk=False) else: @@ -134,7 +125,6 @@ async def websocket(ws: WebSocket, audio_duration[client_id] += audio_to_asr[client_id][-1].duration_seconds logger.debug(asr_result_words) - # Копим ответы для пунктуации ws_collected_asr_res[client_id][f"channel_{1}"].append(asr_result_words) except Exception as e: @@ -142,7 +132,7 @@ async def websocket(ws: WebSocket, else: if len(asr_result_words.get("data").get("text")) == 0 or asr_result_words.get("data").get("text") == ' ': if wait_null_answers: - if not await send_messages(ws, _silence = True, _data = None, _error = None, _channel_name=channel_name): + if not await send_messages(ws, _silence=True, _data=None, _error=None, _channel_name=channel_name): logger.error(f"send_message not ok work canceled") try: del audio_overlap[client_id] @@ -153,12 +143,11 @@ async def websocket(ws: WebSocket, except Exception as e: logger.error(f"error clearing globals after abnormal closing socket - {e}") return - # await asyncio.sleep(0.01) else: logger.debug("sending silence partials skipped") continue else: - if not await send_messages(ws, _silence=False, _data=asr_result_words, _error=None, _channel_name = channel_name): + if not await send_messages(ws, _silence=False, _data=asr_result_words, _error=None, _channel_name=channel_name): logger.error(f"send_message not ok work canceled") try: del audio_overlap[client_id] @@ -189,8 +178,6 @@ async def websocket(ws: WebSocket, logger.error(f"error clearing globals after abnormal closing socket - {e} in channel {channel_name}") return - # Передаём на распознавание собранный не полный буфер - # перевод в семплы для распознавания. audio_to_asr[client_id].append(audio_overlap[client_id] + audio_buffer[client_id]) logger.debug(f'итоговое сообщение - {audio_to_asr[client_id][-1].duration_seconds} секунд') @@ -207,10 +194,8 @@ async def websocket(ws: WebSocket, last_result = process_single_token_vocab_output(last_asr_result_w_conf, audio_duration[client_id]) logger.debug(f'Последний результат {last_result.get("data").get("text")} in channel {channel_name}') - # Добавляем ответ для пунктуации ws_collected_asr_res[client_id][f"channel_{1}"].append(last_result) - except Exception as e: logger.error(f"last_asr_result_w_conf error - {e}") @@ -232,7 +217,6 @@ async def websocket(ws: WebSocket, logger.error(f"await do_sensitizing - {e}") error_description = f"do_sensitizing - {e}" - # if not await send_messages(ws, _silence=is_silence, _data=last_result, _error=error_description, _last_message=True, _sentenced_data=sentenced_data, _channel_name=channel_name): logger.error(f"send_message not ok work canceled in channel {channel_name}") diff --git a/api/v1/endpoints/health.py b/api/v1/endpoints/health.py new file mode 100644 index 0000000..0607294 --- /dev/null +++ b/api/v1/endpoints/health.py @@ -0,0 +1,59 @@ +import logging +from fastapi import APIRouter +from models.fast_api_models import V1BaseResponse, IsAliveData +from api.legacy.is_alive import get_gpu_free_memory +from utils.pre_start_init import audio_to_asr + +router = APIRouter(prefix="/health", tags=["Health"]) +logger = logging.getLogger(__name__) + + +@router.get("/live", response_model=V1BaseResponse) +async def health_live(): + """Liveness probe для Kubernetes.""" + return V1BaseResponse( + success=True, + error_description=None, + data={"status": "ok"} + ) + + +@router.get("/ready", response_model=V1BaseResponse) +async def health_ready(): + """Readiness probe — базовая заглушка до реализации Задачи 4.2.""" + # TODO: проверить загрузку моделей, доступность GPU, свободную память + return V1BaseResponse( + success=True, + error_description=None, + data={"status": "ok", "checks": {"models_loaded": True, "gpu_available": None}} + ) + + +@router.get("/is_alive", response_model=V1BaseResponse) +async def health_is_alive(): + """Перенос текущего is_alive для консистентности API v1.""" + logger.info('GET /api/v1/health/is_alive') + error_description = None + tasks_in_work = len(audio_to_asr) + error, free_mb, gpu_load, temperature = get_gpu_free_memory() + state = "idle" if tasks_in_work == 0 else "in_work" + + if error: + error_description = error.get("error", None) + return V1BaseResponse( + success=True, + error_description=error_description, + data=None + ) + else: + return V1BaseResponse( + success=True, + error_description=None, + data=IsAliveData( + state=state, + tasks_in_work=tasks_in_work, + free_memory_mb=free_mb, + gpu_load_percent=gpu_load, + temperature_celsius=temperature + ) + ) diff --git a/api/v1/endpoints/root.py b/api/v1/endpoints/root.py new file mode 100644 index 0000000..abd9d0c --- /dev/null +++ b/api/v1/endpoints/root.py @@ -0,0 +1,31 @@ +from fastapi import APIRouter +from models.fast_api_models import V1BaseResponse +from config import settings + +router = APIRouter(prefix="", tags=["System"]) + + +@router.get("/", response_model=V1BaseResponse) +async def root(): + """ + Корневой эндпоинт API v1. + + Returns: + V1BaseResponse: базовый ответ с приветственным сообщением. + """ + return V1BaseResponse( + success=True, + error_description=None, + data={ + "message": "No_service_selected", + "available_endpoints": { + "POST /api/v1/asr/url": "ASR by URL", + "POST /api/v1/asr/file": "ASR by file upload", + "GET /api/v1/health/is_alive": "Service health check", + "WS /api/v1/asr/ws": "WebSocket streaming ASR", + "GET /docs": "API documentation", + "/demo": "DEMO UI page" + }, + "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs" + } + ) diff --git a/main.py b/main.py index a53d6d5..821722b 100644 --- a/main.py +++ b/main.py @@ -20,8 +20,8 @@ from VoiceActivityDetector import vad from routes.ws_audio_transkrib import router as ws_audio_transkrib_router -from routes.v1 import router as v1_router -from routes.legacy import router as legacy_router +from api.legacy import router as legacy_router +from api.v1.api import router as api_v1_router import models from config import WS_DESCRIPTION @@ -161,7 +161,7 @@ async def lifespan(app): # Routers app.include_router(ws_audio_transkrib_router, tags=["legacy"]) app.include_router(legacy_router, tags=["legacy"]) -app.include_router(v1_router, tags=["v1"]) +app.include_router(api_v1_router, tags=["api/v1"]) try: if __name__ == '__main__': diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 275cf22..3ef5df1 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -1,4 +1,4 @@ -from pydantic import BaseModel, HttpUrl, Field, ConfigDict +from pydantic import BaseModel, HttpUrl, Field, ConfigDict, RootModel from typing import Union, Annotated, Optional, Any, List, Dict from fastapi import UploadFile @@ -23,13 +23,11 @@ class V1BaseResponse(BaseModel): """ success: bool = True error_description: Optional[str] = None - data: dict = {} + data: Any = {} -class RawData(BaseModel): - """Структура сырых данных ASR.""" - result: Optional[List[Dict[str, Any]]] = None - text: Optional[str] = None +class RawData(RootModel[Dict[str, Any]]): + """Структура сырых данных ASR. Словарь каналов, где каждый канал — список результатов.""" class SentencedData(BaseModel): diff --git a/routes/__init__.py b/routes/__init__.py deleted file mode 100644 index 934f058..0000000 --- a/routes/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from . import ws_audio_transkrib diff --git a/routes/v1/__init__.py b/routes/v1/__init__.py deleted file mode 100644 index d335a31..0000000 --- a/routes/v1/__init__.py +++ /dev/null @@ -1,11 +0,0 @@ -from fastapi import APIRouter -from .is_alive import router as is_alive_router -from .post_by_url import router as post_by_url_router -from .post_by_file_FORM import router as post_by_file_router -from .root import router as root_router - -router = APIRouter(prefix="/v1") -router.include_router(root_router) -router.include_router(is_alive_router) -router.include_router(post_by_url_router) -router.include_router(post_by_file_router) diff --git a/routes/v1/is_alive.py b/routes/v1/is_alive.py deleted file mode 100644 index 9b17702..0000000 --- a/routes/v1/is_alive.py +++ /dev/null @@ -1,35 +0,0 @@ -from fastapi import APIRouter -import logging -from utils.pre_start_init import audio_to_asr -from routes.legacy.is_alive import get_gpu_free_memory -from models.fast_api_models import V1IsAliveResponse, IsAliveData - -router = APIRouter() - -@router.get("/is_alive", response_model=V1IsAliveResponse) -async def check_if_service_is_alive_v1(): - logging.info('GET /v1/is_alive') - error_description = None - tasks_in_work = len(audio_to_asr) - error, free_mb, gpu_load, temperature = get_gpu_free_memory() - state = "idle" if tasks_in_work == 0 else "in_work" - - if error: - error_description = error.get("error",None) - return V1IsAliveResponse( - success=True, - error_description=error_description, - data=None - ) - else: - return V1IsAliveResponse( - success=True, - error_description=None, - data=IsAliveData( - state=state, - tasks_in_work=tasks_in_work, - free_memory_mb=free_mb, - gpu_load_percent=gpu_load, - temperature_celsius=temperature - ) - ) diff --git a/routes/v1/root.py b/routes/v1/root.py deleted file mode 100644 index ca8a874..0000000 --- a/routes/v1/root.py +++ /dev/null @@ -1,28 +0,0 @@ -from fastapi import APIRouter -from models.fast_api_models import V1BaseResponse -from config import settings -router = APIRouter() - - -@router.get("/", response_model=V1BaseResponse) -async def root(): - """ - Корневой эндпоинт API v1. - - Returns: - V1BaseResponse: базовый ответ с приветственным сообщением. - """ - return V1BaseResponse( - success=True, - error_description=None, - data={"message": "No_service_selected", - "available_endpoints": { - "POST v1/post_one_step_req": "ASR by URL", - "POST v1/post_file": "ASR by file upload", - "GET v1/is_alive": "Service health check", - "WS /ws": "WebSocket streaming ASR", - "GET /docs": "API documentation", - "/demo": "DEMO UI page" - }, - "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs"} - ) diff --git a/templates/index.html b/templates/index.html index 006ef5c..5b6ab03 100644 --- a/templates/index.html +++ b/templates/index.html @@ -194,7 +194,7 @@

Ответы сервера:

serverResponses.innerHTML = '

Загрузка...

'; - fetch('/post_one_step_req', { + fetch('/api/v1/asr/url', { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -255,7 +255,7 @@

Ответы сервера:

async function sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate) { return new Promise((resolve, reject) => { const wsProtocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; - const wsUrl = `${wsProtocol}//${window.location.host}/ws`; + const wsUrl = `${wsProtocol}//${window.location.host}/api/v1/asr/ws`; const socket = new WebSocket(wsUrl); const doDialogue = document.getElementById('ws_do_dialogue').checked; @@ -367,7 +367,7 @@

Ответы сервера:

try { serverResponses.innerHTML = '

Загрузка...

'; - const response = await fetch('/post_file', { + const response = await fetch('/api/v1/asr/file', { method: 'POST', body: formData }); @@ -394,4 +394,4 @@

Ответы сервера:

- \ No newline at end of file + diff --git a/tests/test_http_routes.py b/tests/test_http_routes.py index 120bbda..c73b178 100644 --- a/tests/test_http_routes.py +++ b/tests/test_http_routes.py @@ -7,7 +7,7 @@ from config import settings from models.fast_api_models import V1BaseResponse as BaseResponse -BASE_URL = f"http://127.0.0.1:{settings.PORT}/v1" +BASE_URL = f"http://127.0.0.1:{settings.PORT}/api/v1" async def _assert_base_response(body: dict, expect_success: bool): @@ -27,7 +27,7 @@ async def test_post_by_url_success(): } try: resp = await client.post( - f"{BASE_URL}/post_one_step_req", json=payload, timeout=120.0 + f"{BASE_URL}/asr/url", json=payload, timeout=120.0 ) except httpx.ConnectError: pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") @@ -46,7 +46,7 @@ async def test_post_by_url_validation_error(): payload = {"keep_raw": True} # отсутствует AudioFileUrl try: resp = await client.post( - f"{BASE_URL}/post_one_step_req", json=payload, timeout=10.0 + f"{BASE_URL}/asr/url", json=payload, timeout=10.0 ) except httpx.ConnectError: pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") @@ -75,7 +75,7 @@ async def test_post_by_file_success(): } try: resp = await client.post( - f"{BASE_URL}/post_file", data=data, files=files, timeout=20.0 + f"{BASE_URL}/asr/file", data=data, files=files, timeout=20.0 ) except httpx.ConnectError: pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") @@ -94,7 +94,7 @@ async def test_post_by_file_validation_error(): data = {"keep_raw": "true"} try: resp = await client.post( - f"{BASE_URL}/post_file", data=data, timeout=10.0 + f"{BASE_URL}/asr/file", data=data, timeout=10.0 ) except httpx.ConnectError: pytest.fail(f"Сервер не отвечает по адресу {BASE_URL}. Убедитесь, что приложение запущено.") From 9f5c78e8d7f461b98a306e329ab0af5ce3782ad5 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Tue, 5 May 2026 16:16:53 +0300 Subject: [PATCH 25/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=BE=20"Deprecated"=20=D0=B2=20legacy=20=D1=80?= =?UTF-8?q?=D0=BE=D1=83=D1=82=D1=8B.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/legacy/__init__.py | 6 + api/legacy/demo_page.py | 5 +- api/legacy/is_alive.py | 5 +- api/legacy/post_by_file_FORM.py | 4 +- api/legacy/post_by_url.py | 4 + api/legacy/root.py | 5 + main.py | 32 ++++ routes/__init__.py | 1 + routes/ws_audio_transkrib.py | 263 ++++++++++++++++++++++++++++++++ 9 files changed, 322 insertions(+), 3 deletions(-) create mode 100644 routes/__init__.py create mode 100644 routes/ws_audio_transkrib.py diff --git a/api/legacy/__init__.py b/api/legacy/__init__.py index 4e13035..ba1557b 100644 --- a/api/legacy/__init__.py +++ b/api/legacy/__init__.py @@ -1,3 +1,4 @@ +import logging from fastapi import APIRouter from .is_alive import router as is_alive_router from .post_by_url import router as post_by_url_router @@ -5,6 +6,11 @@ from .root import router as root_router from .demo_page import router as demo_page_router +logger = logging.getLogger(__name__) +logger.warning( + "Legacy API routes are deprecated. Use /api/v1/ instead.", +) + router = APIRouter() router.include_router(root_router) router.include_router(demo_page_router) diff --git a/api/legacy/demo_page.py b/api/legacy/demo_page.py index ee5b983..0dcf201 100644 --- a/api/legacy/demo_page.py +++ b/api/legacy/demo_page.py @@ -1,8 +1,8 @@ +import logging from fastapi import APIRouter, Request from fastapi.responses import HTMLResponse from fastapi.templating import Jinja2Templates -import logging logger = logging.getLogger(__name__) router = APIRouter() @@ -12,6 +12,9 @@ @router.get("/demo", response_class=HTMLResponse) async def demo_page(request: Request): + logger.warning( + "Legacy endpoint /demo is deprecated. Use /api/v1/ instead.", + ) return templates.TemplateResponse( "index.html", {"request": request} # Контекст для Jinja2 (если нужно) diff --git a/api/legacy/is_alive.py b/api/legacy/is_alive.py index 3f61eba..ffcdaee 100644 --- a/api/legacy/is_alive.py +++ b/api/legacy/is_alive.py @@ -3,7 +3,7 @@ import pynvml from utils.pre_start_init import audio_to_asr - +logger = logging.getLogger(__name__) router = APIRouter() def get_gpu_free_memory(): @@ -25,6 +25,9 @@ def get_gpu_free_memory(): @router.get("/is_alive") async def check_if_service_is_alive(): + logger.warning( + "Legacy endpoint /is_alive is deprecated. Use /api/v1/health/is_alive instead.", + ) error_description = None logging.info('GET_is_alive') tasks_in_work = len(audio_to_asr) diff --git a/api/legacy/post_by_file_FORM.py b/api/legacy/post_by_file_FORM.py index 83b2c61..79f8bdf 100644 --- a/api/legacy/post_by_file_FORM.py +++ b/api/legacy/post_by_file_FORM.py @@ -11,7 +11,6 @@ from Diarisation import get_diarizer from Diarisation.do_diarize import Diarizer -import logging logger = logging.getLogger(__name__) router = APIRouter() @@ -52,6 +51,9 @@ async def async_receive_file_legacy( punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), diarizer: Diarizer = Depends(get_diarizer) ) -> BaseResponse: + logger.warning( + "Legacy endpoint /post_file is deprecated. Use /api/v1/asr/file instead.", + ) # Сохраняем файл на диск асинхронно try: buffer = BytesIO(await file.read()) diff --git a/api/legacy/post_by_url.py b/api/legacy/post_by_url.py index 472cd8b..5da6069 100644 --- a/api/legacy/post_by_url.py +++ b/api/legacy/post_by_url.py @@ -39,6 +39,10 @@ async def post(params: SyncASRRequest, :param:batch_size: Union[int, None] = Размер пакета для обработки. \n """ + logger.warning( + "Legacy endpoint /post_one_step_req is deprecated. Use /api/v1/asr/url instead.", + ) + # Получаем файл post_id = uuid.uuid4() if params.AudioFileUrl: diff --git a/api/legacy/root.py b/api/legacy/root.py index 084ae1c..76e7bd7 100644 --- a/api/legacy/root.py +++ b/api/legacy/root.py @@ -1,7 +1,9 @@ +import logging from fastapi import APIRouter from models.fast_api_models import V1BaseResponse from config import settings router = APIRouter() +logger = logging.getLogger(__name__) @router.get("/") @@ -12,6 +14,9 @@ async def root(): Returns: V1BaseResponse: базовый ответ с приветственным сообщением. """ + logger.warning( + "Legacy endpoint / is deprecated. Use /api/v1/ instead.", + ) return {"message": "No_service_selected", "available_endpoints": { "POST v1/post_one_step_req": "ASR by URL", diff --git a/main.py b/main.py index 821722b..e4d9068 100644 --- a/main.py +++ b/main.py @@ -46,6 +46,7 @@ async def send_with_request_id(message): if message["type"] == "http.response.start": headers = MutableHeaders(raw=message["headers"]) headers["X-Request-ID"] = request_id + message["headers"] = headers.raw await send(message) try: @@ -54,6 +55,28 @@ async def send_with_request_id(message): request_id_var.reset(token) +class DeprecationHeaderMiddleware: + def __init__(self, app): + self.app = app + + async def __call__(self, scope, receive, send): + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + async def send_with_deprecation(message): + if message["type"] == "http.response.start": + headers = MutableHeaders(raw=message["headers"]) + path = scope.get("path", "") + if path in {"/root/", "/root/demo", "/root/is_alive", "/root/post_file", "/root/post_one_step_req", "/root/ws"}: + headers["Deprecation"] = "true" + headers["Warning"] = f'299 - "Legacy API is deprecated. Use /api/v1/ instead."' + message["headers"] = headers.raw + await send(message) + + await self.app(scope, receive, send_with_deprecation) + + @asynccontextmanager async def lifespan(app): # Настройка логирования до любых других операций @@ -122,9 +145,18 @@ async def lifespan(app): description=WS_DESCRIPTION ) +# Deprecation warning for legacy endpoints at startup +logger.warning( + "Legacy endpoints (/ws, /post_file, /post_one_step_req, /is_alive, /demo, /) are deprecated. " + "Use /api/v1/ instead.", +) + # RequestID middleware app.add_middleware(RequestIDMiddleware) +# Deprecation header middleware for legacy endpoints +app.add_middleware(DeprecationHeaderMiddleware) + # ProxyHeaders middleware app.add_middleware( ProxyHeadersMiddleware, diff --git a/routes/__init__.py b/routes/__init__.py new file mode 100644 index 0000000..934f058 --- /dev/null +++ b/routes/__init__.py @@ -0,0 +1 @@ +from . import ws_audio_transkrib diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py new file mode 100644 index 0000000..6c41364 --- /dev/null +++ b/routes/ws_audio_transkrib.py @@ -0,0 +1,263 @@ +from pydub import AudioSegment + +import ujson +import logging +from config import settings +import uuid +from io import BytesIO + +from fastapi import APIRouter, WebSocket +from utils.chunk_doing import find_last_speech_position +from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration,ws_collected_asr_res +from utils.send_messages import send_messages +from utils.tokens_to_Result import process_single_token_vocab_output +from utils.resamppling import async_resample_audiosegment + +from fastapi import Depends +from Recognizer import get_recognizer, Recognizer + +from Recognizer.engine.sentensizer import do_sensitizing +from Recognizer.engine.stream_recognition import simple_recognise +from Punctuation import get_punctuator, SbertPuncCaseOnnx + +router = APIRouter() +logger = logging.getLogger(__name__) + + +@router.websocket("/ws") +async def websocket(ws: WebSocket, + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + ): + wait_null_answers=True + client_id = uuid.uuid4() + logger.warning( + "Legacy WebSocket endpoint /ws is deprecated. Use /api/v1/asr/ws instead.", + ) + logger.debug(f'Принят новый сокет id = {client_id}') + audio_buffer[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + audio_overlap[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + audio_duration[client_id] = 0 + audio_to_asr[client_id] = list() + ws_collected_asr_res[client_id] = {f"channel_{1}": list()} + do_dialogue = False + do_punctuation = False + audio_format = 'raw' + sample_rate = settings.BASE_SAMPLE_RATE # Если не получен фреймрейт в конфиге сокета, по попытается принять с конфигом модели. + sentenced_data = None + error_description = None + + await ws.accept() + channel_name = str() + + while True: + try: + message = await ws.receive() + except Exception as wse: + logger.error(f"receive WebSocketException - {wse}") + return + + if isinstance(message, dict) and message.get('text'): + try: + if message.get('text') and 'config' in message.get('text'): + json_cfg = ujson.loads(message.get('text'))['config'] + audio_format = json_cfg.get("audio_format", 'pcm16') + sample_rate = json_cfg.get('sample_rate') + wait_null_answers = json_cfg.get('wait_null_answers', wait_null_answers) + do_dialogue = json_cfg.get("do_dialogue", False) + do_punctuation = json_cfg.get("do_punctuation", False) + try: + channel_name = message.get('text').get("channelName") + except Exception as e: + channel_name = "Null" + logger.debug("ChannelName not parsed") + logger.info(f"Task received, config - {message.get('text')}") + continue + + elif message.get('text') and 'eof' in message.get('text'): + logger.info(f"EOF received in channel {channel_name}") + break + else: + logger.error(f"Can`t recognise text part of message {message.get('text')} in channel {channel_name}") + + except Exception as e: + logger.error(f'Error text message compiling. Message:{message} - error:{e} in channel {channel_name}') + elif isinstance(message, dict) and message.get('bytes'): + try: + # Получаем новый чанк с данными + chunk = message.get('bytes') + + if audio_format == 'pcm16': + # Проверяем и добавляем недостающие нулевые байты в чанки. + if len(chunk) % 2 != 0: + chunk += bytes(2 - (len(chunk) % 2)) + + # Переводим чанк в объект Audiosegment + audiosegment_chunk = AudioSegment( + chunk, + frame_rate = sample_rate, # Частота дискретизации + sample_width = 2, # Ширина сэмпла (2 байта для int16) + channels = 1 # Количество каналов. По умолчанию - 1, Моно. + ) + + else: + try: + buffer = BytesIO(chunk) + buffer.seek(0) + # buffer.write() + audiosegment_chunk = AudioSegment.from_file(buffer) + + except Exception as e: + logger.error(f"Ошибка принятия аудио - {e} in channel {channel_name}") + else: + logger.debug(f"Чанк принят и распознан in channel {channel_name}") + + # Приводим фреймрейт к фреймрейту модели + if audiosegment_chunk.frame_rate != settings.BASE_SAMPLE_RATE: + audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, settings.BASE_SAMPLE_RATE) + + if audiosegment_chunk.channels != 1: + audiosegment_chunk = audiosegment_chunk.set_channels(1) + # Копим буфер + audio_buffer[client_id] += audiosegment_chunk + + # Накопили больше нормы + if (audio_overlap[client_id]+audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: + # Проверяем новый чанк перед объединением (там же режем хвост и добавляем его при необходимости) + await find_last_speech_position(client_id, is_last_chunk=False) + + else: + continue + except Exception as e: + logger.error(f"AcceptWaveform error - {e} in channel {channel_name}") + else: + try: + asr_result = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) + asr_result_words = process_single_token_vocab_output(asr_result, audio_duration[client_id]) + audio_duration[client_id] += audio_to_asr[client_id][-1].duration_seconds + logger.debug(asr_result_words) + + # Копим ответы для пунктуации + ws_collected_asr_res[client_id][f"channel_{1}"].append(asr_result_words) + + except Exception as e: + logger.error(f"recognizer.get_result(stream()) error - {e}") + else: + if len(asr_result_words.get("data").get("text")) == 0 or asr_result_words.get("data").get("text") == ' ': + if wait_null_answers: + if not await send_messages(ws, _silence = True, _data = None, _error = None, _channel_name=channel_name): + logger.error(f"send_message not ok work canceled") + try: + del audio_overlap[client_id] + del audio_buffer[client_id] + del audio_to_asr[client_id] + del audio_duration[client_id] + del ws_collected_asr_res[client_id] + except Exception as e: + logger.error(f"error clearing globals after abnormal closing socket - {e}") + return + # await asyncio.sleep(0.01) + else: + logger.debug("sending silence partials skipped") + continue + else: + if not await send_messages(ws, _silence=False, _data=asr_result_words, _error=None, _channel_name = channel_name): + logger.error(f"send_message not ok work canceled") + try: + del audio_overlap[client_id] + del audio_buffer[client_id] + del audio_to_asr[client_id] + del audio_duration[client_id] + del ws_collected_asr_res[client_id] + except Exception as e: + logger.error(f"error clearing globals after abnormal closing socket - {e}") + return + elif isinstance(message, dict) and message.get('type') == "websocket.disconnect": + description = f"Channel {channel_name} closed from outside" + logger.error(description) + break + else: + error_description = f"Can`t parse message - {message} in channel {channel_name}" + logger.error(error_description) + + if not await send_messages(ws, _silence=False, _data=None, _error=error_description, _channel_name=channel_name): + logger.error(f"send_message not ok work canceled in channel {channel_name}") + try: + del audio_overlap[client_id] + del audio_buffer[client_id] + del audio_to_asr[client_id] + del audio_duration[client_id] + del ws_collected_asr_res[client_id] + except Exception as e: + logger.error(f"error clearing globals after abnormal closing socket - {e} in channel {channel_name}") + return + + # Передаём на распознавание собранный не полный буфер + # перевод в семплы для распознавания. + audio_to_asr[client_id].append(audio_overlap[client_id] + audio_buffer[client_id]) + logger.debug(f'итоговое сообщение - {audio_to_asr[client_id][-1].duration_seconds} секунд') + + try: + try: + if audio_to_asr[client_id][-1].duration_seconds < 2: + audio_to_asr[client_id][-1] = audio_to_asr[client_id][-1] + AudioSegment.silent(1000, frame_rate=sample_rate) + except Exception as e: + logger.error(f"Ошибка дополнения тишиной последнего чанка - {e} in channel {channel_name}") + last_result = None + error_description = f"Ошибка дополнения тишиной последнего чанка - {e} in channel {channel_name}" + else: + last_asr_result_w_conf = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) + last_result = process_single_token_vocab_output(last_asr_result_w_conf, audio_duration[client_id]) + logger.debug(f'Последний результат {last_result.get("data").get("text")} in channel {channel_name}') + + # Добавляем ответ для пунктуации + ws_collected_asr_res[client_id][f"channel_{1}"].append(last_result) + + + except Exception as e: + logger.error(f"last_asr_result_w_conf error - {e}") + + else: + if len(last_result.get("data").get("text")) == 0: + is_silence = True + last_result = None + elif last_result.get("data").get("text") == ' ': + is_silence = True + last_result = None + else: + logger.debug(last_result) + is_silence = False + + if do_dialogue: + try: + sentenced_data = await do_sensitizing(ws_collected_asr_res[client_id], do_punctuation, punctuator=punctuator) + except Exception as e: + logger.error(f"await do_sensitizing - {e}") + error_description = f"do_sensitizing - {e}" + + # + if not await send_messages(ws, _silence=is_silence, _data=last_result, _error=error_description, _last_message=True, + _sentenced_data=sentenced_data, _channel_name=channel_name): + logger.error(f"send_message not ok work canceled in channel {channel_name}") + try: + del audio_overlap[client_id] + del audio_buffer[client_id] + del audio_to_asr[client_id] + del audio_duration[client_id] + del ws_collected_asr_res[client_id] + except Exception as e: + logger.error(f"error clearing globals after abnormal closing socket - {e} in channel {channel_name}") + return + + logger.info(f"Closing connection {channel_name}") + await ws.close() + + try: + del audio_overlap[client_id] + del audio_buffer[client_id] + del audio_to_asr[client_id] + del audio_duration[client_id] + del ws_collected_asr_res[client_id] + except Exception as e: + logger.error(f"error clearing globals after NORMAL closing socket - {e} in channel {channel_name}") + return From cd4846843039d8f021cd045e9850bb08abc6504d Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Wed, 6 May 2026 11:38:21 +0300 Subject: [PATCH 26/46] =?UTF-8?q?=D0=93=D0=BE=D1=82=D0=BE=D0=B2=D0=B8?= =?UTF-8?q?=D0=BC=20=D0=BA=20=D0=B2=D1=8B=D0=B5=D0=B7=D0=B4=D1=83=20=D0=B1?= =?UTF-8?q?=D0=B8=D0=B7=D0=BD=D0=B5=D1=81-=D0=BB=D0=BE=D0=B3=D0=B8=D0=BA?= =?UTF-8?q?=D0=B8=20WS=20=D0=B8=D0=B7=20=D1=80=D0=BE=D1=83=D1=82=D0=B0.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.py | 23 +++- main.py | 2 + models/ws_models.py | 150 +++++++++++++++++++++ routes/ws_audio_transkrib.py | 2 + services/__init__.py | 1 + services/ws_handler.py | 251 ++++++++++++++++++++++++++++++++++ services/ws_manager.py | 191 ++++++++++++++++++++++++++ services/ws_metrics.py | 169 +++++++++++++++++++++++ services/ws_session.py | 145 ++++++++++++++++++++ tests/test_ws_handler.py | 212 +++++++++++++++++++++++++++++ tests/test_ws_manager.py | 143 ++++++++++++++++++++ tests/test_ws_metrics.py | 106 +++++++++++++++ tests/test_ws_models.py | 253 +++++++++++++++++++++++++++++++++++ tests/test_ws_session.py | 114 ++++++++++++++++ 14 files changed, 1759 insertions(+), 3 deletions(-) create mode 100644 models/ws_models.py create mode 100644 services/__init__.py create mode 100644 services/ws_handler.py create mode 100644 services/ws_manager.py create mode 100644 services/ws_metrics.py create mode 100644 services/ws_session.py create mode 100644 tests/test_ws_handler.py create mode 100644 tests/test_ws_manager.py create mode 100644 tests/test_ws_metrics.py create mode 100644 tests/test_ws_models.py create mode 100644 tests/test_ws_session.py diff --git a/config.py b/config.py index f2d50ff..6495de4 100644 --- a/config.py +++ b/config.py @@ -114,6 +114,17 @@ class Settings(BaseSettings): # Quota settings GUEST_DAILY_QUOTA: int = 10 + # WebSocket settings + WS_MAX_CONNECTIONS: int = 100 + WS_MAX_BUFFER_DURATION_SEC: float = 300.0 + WS_IDLE_TIMEOUT_SEC: float = 60.0 + WS_PING_TIMEOUT_SEC: float = 20.0 + WS_MAX_MESSAGE_SIZE_MB: float = 10.0 + WS_MAX_SESSION_DURATION_SEC: float = 300.0 + WS_STATUS_BROADCAST_INTERVAL_SEC: float = 5.0 + WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT: float = 90.0 + WS_STATUS_BUSY_CONNECTIONS_THRESHOLD_PCT: float = 80.0 + @field_validator( 'IS_PROD', 'MAKE_MONO', 'USE_BATCH', 'VAD_WITH_GPU', 'CAN_PUNCTUATE', 'PUNCTUATE_WITH_GPU', 'CAN_DIAR', @@ -218,16 +229,22 @@ def _compute_derived(self): } ``` -### Пример передачи данных. +### Пример передачи данных (JSON + base64) -Периодически отправляйте raw_audio_data - PCM, 16-bit, mono +Используется при `audio_transport: "json_base64"` (по умолчанию): ```json { - "bytes": binary + "type": "audio_chunk", + "audio_base64": "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAEAfAAABAAgAZGF0YQAAAAA=", + "seq_num": 0 } ``` +### Пример передачи данных (Binary) + +Используется при `audio_transport: "binary"`. Отправляйте WebSocket **binary frame** напрямую (без JSON-обёртки). Сервер читает его через `receive_bytes()`. + ### Пример EOF По завершении отправьте: diff --git a/main.py b/main.py index e4d9068..8b7a126 100644 --- a/main.py +++ b/main.py @@ -1,4 +1,5 @@ import logging +import time import uvicorn from config import settings import os @@ -87,6 +88,7 @@ async def lifespan(app): # on_start logger.debug("Приложение FastAPI запущено") + app.state.start_time = time.time() # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) diff --git a/models/ws_models.py b/models/ws_models.py new file mode 100644 index 0000000..8c4044c --- /dev/null +++ b/models/ws_models.py @@ -0,0 +1,150 @@ +import time +import base64 +from enum import Enum +from typing import Literal, Union, Annotated + +from pydantic import BaseModel, Field, TypeAdapter + + +class WSMessageType(str, Enum): + config = "config" + audio_chunk = "audio_chunk" + ping = "ping" + pong = "pong" + status_request = "status_request" + status_response = "status_response" + error = "error" + eos = "eos" + partial_result = "partial_result" + final_result = "final_result" + + +class WSBaseMessage(BaseModel): + type: WSMessageType + timestamp: float | None = Field(default_factory=time.time) + + +class WSConfigMessage(WSBaseMessage): + type: Literal[WSMessageType.config] = WSMessageType.config + sample_rate: int = Field(default=16000, ge=8000, le=48000) + language: str = "ru" + wait_null_answers: bool = True + enable_diarization: bool = False + num_speakers: int = Field(default=-1, ge=-1, le=10) + enable_punctuation: bool = False + audio_format: str = "pcm16" + audio_transport: Literal["json_base64", "binary"] = Field( + default="json_base64", + description='Транспорт аудио: "json_base64" — аудио внутри JSON как base64-строка; ' + '"binary" — сырые байты через WebSocket binary frames (receive_bytes()).' + ) + channel_name: str | None = None + + +class WSAudioMessage(WSBaseMessage): + type: Literal[WSMessageType.audio_chunk] = WSMessageType.audio_chunk + audio_base64: str | None = None + seq_num: int = Field(default=0, ge=0) + + +class WSPingMessage(WSBaseMessage): + type: Literal[WSMessageType.ping] = WSMessageType.ping + + +class WSPongMessage(WSBaseMessage): + type: Literal[WSMessageType.pong] = WSMessageType.pong + + +class WSStatusRequest(WSBaseMessage): + type: Literal[WSMessageType.status_request] = WSMessageType.status_request + command: str = "get_status" + + +class WSStatusResponse(WSBaseMessage): + type: Literal[WSMessageType.status_response] = WSMessageType.status_response + adapter_status: Literal["idle", "busy", "overloaded"] = "idle" + gpu_memory_free_mb: int | None = None + gpu_memory_total_mb: int | None = None + cpu_memory_free_mb: int | None = None + cpu_memory_total_mb: int | None = None + active_tasks_count: int = 0 + active_connections_count: int = 0 + queue_depth: int = 0 + uptime_sec: float = 0.0 + + +class WSWordItem(BaseModel): + conf: float + start: float + end: float + word: str + + +class WSRecognitionData(BaseModel): + result: list[WSWordItem] = Field(default_factory=list) + text: str = "" + + +class WSResultMessage(WSBaseMessage): + type: Literal[WSMessageType.partial_result, WSMessageType.final_result] = WSMessageType.partial_result + channel_name: str = "Null" + silence: bool = False + data: WSRecognitionData + error: str | None = None + last_message: bool = False + sentenced_data: dict | None = None + + +class WSPhraseResult(WSBaseMessage): + type: Literal[WSMessageType.partial_result, WSMessageType.final_result] = WSMessageType.partial_result + text: str + speaker_id: str | None = None + start_time: float + end_time: float + is_final: bool = False + + +class WSErrorMessage(WSBaseMessage): + type: Literal[WSMessageType.error] = WSMessageType.error + code: str = "unknown_error" + message: str + is_fatal: bool = False + + +class WSEosMessage(WSBaseMessage): + type: Literal[WSMessageType.eos] = WSMessageType.eos + + +WSMessage = Annotated[ + Union[ + WSConfigMessage, + WSAudioMessage, + WSPingMessage, + WSPongMessage, + WSStatusRequest, + WSStatusResponse, + WSResultMessage, + WSErrorMessage, + WSEosMessage, + ], + Field(discriminator="type"), +] + + +def wrap_binary_audio(audio_bytes: bytes, seq_num: int = 0) -> WSAudioMessage: + """Оборачивает raw bytes в WSAudioMessage (base64).""" + return WSAudioMessage( + audio_base64=base64.b64encode(audio_bytes).decode("utf-8"), + seq_num=seq_num, + ) + + +# TypeAdapter для дискриминированного union — используется в роутах и тестах +ws_message_adapter = TypeAdapter(WSMessage) + + +def parse_ws_message(raw: str | bytes | dict) -> WSBaseMessage: + """Валидирует входящее WS-сообщение из JSON-строки, bytes или dict.""" + if isinstance(raw, (str, bytes)): + return ws_message_adapter.validate_json(raw) + return ws_message_adapter.validate_python(raw) diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 6c41364..80d52c0 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -20,6 +20,8 @@ from Recognizer.engine.stream_recognition import simple_recognise from Punctuation import get_punctuator, SbertPuncCaseOnnx +#Todo Этот роут должен умереть. + router = APIRouter() logger = logging.getLogger(__name__) diff --git a/services/__init__.py b/services/__init__.py new file mode 100644 index 0000000..0274469 --- /dev/null +++ b/services/__init__.py @@ -0,0 +1 @@ +# services package diff --git a/services/ws_handler.py b/services/ws_handler.py new file mode 100644 index 0000000..c1eb0a2 --- /dev/null +++ b/services/ws_handler.py @@ -0,0 +1,251 @@ +""" +Модуль services/ws_handler.py +Содержит класс MessageRouter и набор стандартных хендлеров +для обработки входящих WebSocket-сообщений в рамках ASR-сессии. +""" + +import base64 +import logging +from typing import Callable, Awaitable, Protocol + +import numpy as np + +from models.ws_models import ( + WSBaseMessage, + WSConfigMessage, + WSAudioMessage, + WSStatusRequest, + WSPingMessage, + WSPongMessage, + WSErrorMessage, + WSStatusResponse, + WSMessageType, +) +from services.ws_session import AudioSession, SessionState +from services.ws_metrics import SystemMetricsCollector + +logger = logging.getLogger(__name__) + + +class WSManagerProtocol(Protocol): + """ + Протокол для менеджера соединений, используемого хендлерами. + + Позволяет MessageRouter и хендлерам работать с любой реализацией + менеджера (ConnectionManager, FakeManager в тестах и т.д.). + """ + + async def send_message(self, client_id: str, message: WSBaseMessage | dict) -> None: + """ + Отправляет сообщение указанному клиенту. + + Args: + client_id: Идентификатор клиента. + message: Сообщение для отправки (Pydantic-модель или dict). + """ + ... + + @property + def active_connections_count(self) -> int: + """Текущее количество активных соединений.""" + ... + + @property + def max_connections(self) -> int: + """Максимально допустимое количество соединений.""" + ... + + +async def handle_config( + message: WSConfigMessage, + session: AudioSession, + manager: WSManagerProtocol, +) -> None: + """ + Обрабатывает сообщение конфигурации от клиента. + + Устанавливает конфигурацию в AudioSession и переводит сессию + в состояние receiving. + + Args: + message: Сообщение конфигурации (WSConfigMessage). + session: Текущая аудио-сессия. + manager: Менеджер соединений (для отправки ответов при необходимости). + """ + session.config = message + session.state = SessionState.receiving + logger.debug( + "Config set for client %s: sample_rate=%d, transport=%s", + session.client_id, + message.sample_rate, + message.audio_transport, + ) + + +async def handle_audio( + message: WSAudioMessage, + session: AudioSession, + manager: WSManagerProtocol, +) -> None: + """ + Обрабатывает аудио-чанк от клиента. + + Декодирует base64 в numpy-массив (int16 -> float32 нормализованный [-1, 1]) + и добавляет во внутренний буфер сессии. + Если буфер переполнен — отправляет клиенту WSErrorMessage. + + Args: + message: Сообщение с аудио (WSAudioMessage). + session: Текущая аудио-сессия. + manager: Менеджер соединений для отправки ошибок клиенту. + """ + if message.audio_base64 is None: + return + + try: + raw_bytes = base64.b64decode(message.audio_base64) + # Преобразуем int16 (PCM16) в float32 нормализованный [-1, 1] + audio_array = np.frombuffer(raw_bytes, dtype=np.int16).astype(np.float32) / 32768.0 + except Exception as exc: + logger.warning( + "Failed to decode audio base64 for client %s: %s", + session.client_id, + exc, + ) + error = WSErrorMessage( + code="decode_error", + message=f"Invalid audio_base64 data: {exc}", + is_fatal=False, + ) + await manager.send_message(session.client_id, error) + return + + added = await session.add_audio(audio_array) + if not added: + logger.warning("Audio buffer overflow for client %s", session.client_id) + error = WSErrorMessage( + code="buffer_overflow", + message="Audio buffer limit exceeded", + is_fatal=False, + ) + await manager.send_message(session.client_id, error) + + +async def handle_status_request( + message: WSStatusRequest, + session: AudioSession, + manager: WSManagerProtocol, + metrics_collector: SystemMetricsCollector, +) -> None: + """ + Обрабатывает запрос статуса от клиента. + + Собирает метрики через SystemMetricsCollector и отправляет + WSStatusResponse запросившему клиенту. + + Args: + message: Запрос статуса (WSStatusRequest). + session: Текущая аудио-сессия. + manager: Менеджер соединений. + metrics_collector: Коллектор системных метрик. + """ + status = metrics_collector.collect( + active_connections=manager.active_connections_count, + max_connections=manager.max_connections, + ) + await manager.send_message(session.client_id, status) + + +async def handle_ping( + message: WSPingMessage, + session: AudioSession, + manager: WSManagerProtocol, +) -> None: + """ + Обрабатывает ping-сообщение. + + Отправляет клиенту pong в ответ. + + Args: + message: Ping-сообщение. + session: Текущая аудио-сессия. + manager: Менеджер соединений. + """ + pong = WSPongMessage() + await manager.send_message(session.client_id, pong) + + +class MessageRouter: + """ + Маршрутизатор входящих WebSocket-сообщений. + + Регистрирует хендлеры для каждого WSMessageType и направляет + сообщения в соответствующий обработчик. + """ + + def __init__(self) -> None: + """ + Инициализирует пустой реестр хендлеров. + """ + self._handlers: dict[WSMessageType, Callable[..., Awaitable[None]]] = {} + + def register_handler( + self, + msg_type: WSMessageType, + handler: Callable[..., Awaitable[None]], + ) -> None: + """ + Регистрирует хендлер для указанного типа сообщения. + + Args: + msg_type: Тип WebSocket-сообщения. + handler: Асинхронная функция-обработчик. + """ + self._handlers[msg_type] = handler + logger.debug("Registered handler for %s", msg_type) + + async def route( + self, + message: WSBaseMessage, + session: AudioSession, + manager: WSManagerProtocol, + metrics_collector: SystemMetricsCollector | None = None, + ) -> None: + """ + Направляет сообщение в зарегистрированный хендлер. + + Если хендлер не найден — отправляет клиенту WSErrorMessage. + Для status_request требуется metrics_collector. + + Args: + message: Входящее сообщение (любая модель WSBaseMessage). + session: Текущая аудио-сессия. + manager: Менеджер соединений. + metrics_collector: Коллектор метрик (опционально, нужен для status_request). + """ + handler = self._handlers.get(message.type) + if handler is None: + logger.warning("No handler registered for message type %s", message.type) + error = WSErrorMessage( + code="unsupported_type", + message=f"Unsupported message type: {message.type}", + is_fatal=False, + ) + await manager.send_message(session.client_id, error) + return + + try: + if message.type == WSMessageType.status_request: + if metrics_collector is None: + raise RuntimeError("metrics_collector required for status_request") + await handler(message, session, manager, metrics_collector) + else: + await handler(message, session, manager) + except Exception as exc: + logger.exception("Handler error for %s: %s", message.type, exc) + error = WSErrorMessage( + code="handler_error", + message=f"Internal handler error: {exc}", + is_fatal=False, + ) + await manager.send_message(session.client_id, error) diff --git a/services/ws_manager.py b/services/ws_manager.py new file mode 100644 index 0000000..01a3f52 --- /dev/null +++ b/services/ws_manager.py @@ -0,0 +1,191 @@ +""" +Модуль services/ws_manager.py +Содержит класс ConnectionManager для централизованного управления +WebSocket-соединениями: подключение, отключение, отправка сообщений, +broadcast статуса и graceful disconnect_all. +""" + +import logging +import asyncio +from typing import Optional, Dict, Any +from dataclasses import dataclass, field + +from fastapi import WebSocket, WebSocketDisconnect + +from models.ws_models import WSBaseMessage, WSStatusResponse + +logger = logging.getLogger(__name__) + + +@dataclass +class ConnectionMeta: + """ + Мета-информация о WebSocket-соединении. + + Attributes: + connected_at: Unix-timestamp установки соединения. + last_activity_at: Unix-timestamp последней активности. + bytes_received: Количество полученных байт. + messages_received: Количество полученных сообщений. + client_ip: IP-адрес клиента (или None). + user_agent: User-Agent клиента (или None). + subscribe_status: Флаг подписки на периодические status_response. + """ + connected_at: float = field(default_factory=lambda: __import__("time").time()) + last_activity_at: float = field(default_factory=lambda: __import__("time").time()) + bytes_received: int = 0 + messages_received: int = 0 + client_ip: Optional[str] = None + user_agent: Optional[str] = None + subscribe_status: bool = False + + +class ConnectionManager: + """ + Централизованный менеджер активных WebSocket-соединений. + + Хранит только объекты WebSocket и мета-информацию в памяти текущего процесса. + В будущем мета может выноситься в StateStore (Redis) для кластеризации. + + Attributes: + max_connections: Максимальное количество одновременных соединений. + active_connections: Словарь client_id -> WebSocket. + connection_meta: Словарь client_id -> ConnectionMeta. + """ + + def __init__(self, max_connections: int = 100) -> None: + """ + Инициализирует менеджер с заданным лимитом соединений. + + Args: + max_connections: Максимально допустимое число активных соединений. + """ + self.max_connections: int = max_connections + self.active_connections: Dict[str, WebSocket] = {} + self.connection_meta: Dict[str, ConnectionMeta] = {} + + @property + def active_connections_count(self) -> int: + """Текущее количество активных соединений.""" + return len(self.active_connections) + + async def connect(self, websocket: WebSocket, client_id: str) -> bool: + """ + Принимает новое WebSocket-соединение. + + Args: + websocket: Объект WebSocket из FastAPI. + client_id: Уникальный идентификатор клиента. + + Returns: + True — соединение установлено и добавлено в реестр. + False — превышен лимит соединений, websocket.close(1008) вызван. + """ + if self.active_connections_count >= self.max_connections: + logger.warning("Max connections (%d) reached, rejecting %s", self.max_connections, client_id) + await websocket.close(code=1008, reason="Server overloaded") + return False + + await websocket.accept() + self.active_connections[client_id] = websocket + self.connection_meta[client_id] = ConnectionMeta() + logger.debug("Client %s connected. Total: %d", client_id, self.active_connections_count) + return True + + async def disconnect(self, client_id: str) -> None: + """ + Закрывает соединение и удаляет клиента из реестров. + + Args: + client_id: Идентификатор клиента для отключения. + """ + websocket = self.active_connections.pop(client_id, None) + self.connection_meta.pop(client_id, None) + if websocket is not None: + try: + await websocket.close() + except Exception as exc: + logger.debug("Error closing websocket for %s: %s", client_id, exc) + logger.debug("Client %s disconnected. Total: %d", client_id, self.active_connections_count) + + async def send_message(self, client_id: str, message: WSBaseMessage | dict | str) -> None: + """ + Отправляет сообщение указанному клиенту. + + Args: + client_id: Идентификатор клиента. + message: Сообщение (Pydantic-модель, dict или строка). + """ + websocket = self.active_connections.get(client_id) + if websocket is None: + logger.warning("Cannot send message: client %s not found", client_id) + return + + if isinstance(message, WSBaseMessage): + payload = message.model_dump_json() + elif isinstance(message, dict): + import json + payload = json.dumps(message, separators=(',', ':')) + else: + payload = str(message) + + try: + await websocket.send_text(payload) + meta = self.connection_meta.get(client_id) + if meta is not None: + meta.last_activity_at = __import__("time").time() + except Exception as exc: + logger.warning("Failed to send message to %s: %s", client_id, exc) + + async def broadcast_status(self, status: WSStatusResponse) -> None: + """ + Отправляет статус всем активным соединениям, подписанным на статус. + + Args: + status: Сообщение со статусом адаптера. + """ + payload = status.model_dump_json() + tasks = [] + for client_id, meta in list(self.connection_meta.items()): + if meta.subscribe_status: + ws = self.active_connections.get(client_id) + if ws is not None: + tasks.append(self._send_text_safe(ws, payload, client_id)) + + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + async def _send_text_safe(self, websocket: WebSocket, payload: str, client_id: str) -> None: + """Внутренний хелпер для безопасной отправки текста.""" + try: + await websocket.send_text(payload) + except Exception as exc: + logger.debug("broadcast_status failed for %s: %s", client_id, exc) + + async def disconnect_all(self, code: int = 1001, reason: str = "Server shutdown") -> None: + """ + Принудительно закрывает все активные соединения. + + Используется при graceful shutdown. + + Args: + code: Код закрытия WebSocket (по умолчанию 1001 — going away). + reason: Причина закрытия. + """ + logger.info("Disconnecting all %d connections", self.active_connections_count) + tasks = [] + for client_id, websocket in list(self.active_connections.items()): + tasks.append(self._close_safe(websocket, code, reason, client_id)) + + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + self.active_connections.clear() + self.connection_meta.clear() + + async def _close_safe(self, websocket: WebSocket, code: int, reason: str, client_id: str) -> None: + """Внутренний хелпер для безопасного закрытия websocket.""" + try: + await websocket.close(code=code, reason=reason) + except Exception as exc: + logger.debug("disconnect_all close error for %s: %s", client_id, exc) diff --git a/services/ws_metrics.py b/services/ws_metrics.py new file mode 100644 index 0000000..8a3222c --- /dev/null +++ b/services/ws_metrics.py @@ -0,0 +1,169 @@ +""" +Модуль services/ws_metrics.py +Содержит класс SystemMetricsCollector для сбора системных метрик +(GPU, CPU, активные задачи, статус адаптера) и формирования +WSStatusResponse для передачи через WebSocket. +""" + +import time +import logging +from typing import Optional, Tuple + +from models.ws_models import WSStatusResponse +from config import settings + +logger = logging.getLogger(__name__) + + +class SystemMetricsCollector: + """ + Собирает системные метрики и вычисляет статус адаптера ASR. + + Attributes: + start_time: Unix-timestamp запуска коллектора (для uptime). + _active_tasks: Счётчик активных задач распознавания. + _nvml_handle: Опциональный handle pynvml (из app.state). + """ + + def __init__( + self, + nvml_handle: Optional[object] = None, + start_time: Optional[float] = None, + ) -> None: + """ + Инициализирует коллектор метрик. + + Args: + nvml_handle: Handle pynvml (например, из app.state.nvml_handle). + start_time: Unix-timestamp старта приложения. По умолчанию — текущее время. + """ + self.start_time: float = start_time if start_time is not None else time.time() + self._nvml_handle: Optional[object] = nvml_handle + self._active_tasks: int = 0 + + def get_gpu_stats(self) -> Tuple[Optional[int], Optional[int]]: + """ + Возвращает свободную и общую память GPU в мегабайтах. + + Returns: + Кортеж (free_mb, total_mb). Если pynvml недоступен или не инициализирован — + возвращает (None, None). + """ + if self._nvml_handle is None: + return None, None + try: + import pynvml + info = pynvml.nvmlDeviceGetMemoryInfo(self._nvml_handle) + free_mb = int(info.free / 1024 / 1024) + total_mb = int(info.total / 1024 / 1024) + return free_mb, total_mb + except Exception as exc: + logger.warning("Failed to get GPU stats: %s", exc) + return None, None + + def get_cpu_stats(self) -> Tuple[Optional[int], Optional[int]]: + """ + Возвращает свободную и общую память RAM в мегабайтах. + + Returns: + Кортеж (free_mb, total_mb). Если psutil недоступен — (None, None). + """ + try: + import psutil + mem = psutil.virtual_memory() + free_mb = int(mem.available / 1024 / 1024) + total_mb = int(mem.total / 1024 / 1024) + return free_mb, total_mb + except Exception: + return None, None + + def increment_tasks(self) -> None: + """Увеличивает счётчик активных задач распознавания на 1.""" + self._active_tasks += 1 + + def decrement_tasks(self) -> None: + """Уменьшает счётчик активных задач на 1 (не ниже 0).""" + self._active_tasks = max(0, self._active_tasks - 1) + + def get_active_tasks_count(self) -> int: + """Возвращает текущее количество активных задач распознавания.""" + return self._active_tasks + + def get_queue_depth(self) -> int: + """ + Возвращает глубину очереди задач на обработку. + + Returns: + Заглушка (0) до реализации очереди в последующих этапах. + """ + return 0 + + def get_adapter_status( + self, + active_connections: int, + max_connections: int, + ) -> str: + """ + Вычисляет статус адаптера: idle / busy / overloaded. + + Пороги берутся из конфигурации: + - GPU memory usage > WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT -> overloaded. + - active_connections / max_connections > WS_STATUS_BUSY_CONNECTIONS_THRESHOLD_PCT -> busy. + + Args: + active_connections: Текущее число активных WebSocket-соединений. + max_connections: Максимально допустимое число соединений. + + Returns: + Строка-статус: "idle", "busy" или "overloaded". + """ + gpu_free, gpu_total = self.get_gpu_stats() + if gpu_total and gpu_total > 0: + gpu_used_pct = (gpu_total - gpu_free) / gpu_total * 100 + threshold = getattr(settings, "WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT", 90.0) + if gpu_used_pct >= threshold: + return "overloaded" + + if max_connections > 0: + conn_ratio = active_connections / max_connections + busy_threshold = getattr(settings, "WS_STATUS_BUSY_CONNECTIONS_THRESHOLD_PCT", 80.0) / 100 + if conn_ratio >= busy_threshold: + return "busy" + + if self._active_tasks > 0: + return "busy" + + return "idle" + + def collect( + self, + active_connections: int = 0, + max_connections: int = 100, + ) -> WSStatusResponse: + """ + Собирает полный набор метрик и возвращает WSStatusResponse. + + Args: + active_connections: Текущее количество активных WS-соединений. + max_connections: Максимально допустимое количество соединений. + + Returns: + WSStatusResponse с актуальными метриками. + """ + gpu_free, gpu_total = self.get_gpu_stats() + cpu_free, cpu_total = self.get_cpu_stats() + uptime = time.time() - self.start_time + + status = self.get_adapter_status(active_connections, max_connections) + + return WSStatusResponse( + adapter_status=status, + gpu_memory_free_mb=gpu_free, + gpu_memory_total_mb=gpu_total, + cpu_memory_free_mb=cpu_free, + cpu_memory_total_mb=cpu_total, + active_tasks_count=self._active_tasks, + active_connections_count=active_connections, + queue_depth=self.get_queue_depth(), + uptime_sec=round(uptime, 2), + ) diff --git a/services/ws_session.py b/services/ws_session.py new file mode 100644 index 0000000..6a40a72 --- /dev/null +++ b/services/ws_session.py @@ -0,0 +1,145 @@ +""" +Модуль services/ws_session.py +Содержит класс AudioSession для управления состоянием и буфером аудио +в рамках одной WebSocket-сессии распознавания речи (ASR). +""" + +import time +from collections import deque +from enum import Enum +from typing import Optional + +import numpy as np + +from models.ws_models import WSConfigMessage +from config import settings + + +class SessionState(str, Enum): + """Состояния жизненного цикла аудио-сессии.""" + connecting = "connecting" + receiving = "receiving" + processing = "processing" + completed = "completed" + error = "error" + + +class AudioSession: + """ + Управляет буфером аудио и состоянием для одного WebSocket-клиента. + + Attributes: + client_id: Уникальный идентификатор сессии. + state: Текущее состояние сессии (SessionState). + buffer: Очередь (deque) numpy-массивов с фрагментами аудио. + config: Конфигурация клиента (WSConfigMessage) или None. + max_buffer_duration_sec: Максимальная суммарная длительность буфера в секундах. + last_activity: Unix-timestamp последней активности (добавления аудио или конфига). + """ + + def __init__( + self, + client_id: str, + max_buffer_duration_sec: Optional[float] = None, + ) -> None: + """ + Инициализирует новую аудио-сессию. + + Args: + client_id: UUID или строковый идентификатор соединения. + max_buffer_duration_sec: Максимальная длительность буфера (сек). + По умолчанию берётся из settings.WS_MAX_BUFFER_DURATION_SEC. + """ + self.client_id: str = client_id + self.state: SessionState = SessionState.connecting + self.buffer: deque[np.ndarray] = deque() + self.config: Optional[WSConfigMessage] = None + self.max_buffer_duration_sec: float = ( + max_buffer_duration_sec + if max_buffer_duration_sec is not None + else getattr(settings, "WS_MAX_BUFFER_DURATION_SEC", 300.0) + ) + self.last_activity: float = time.time() + + @property + def current_buffer_duration_sec(self) -> float: + """ + Вычисляет суммарную длительность аудио в буфере (в секундах). + + Использует sample_rate из config (по умолчанию 16000 Гц), + считая, что каждый элемент буфера — одномерный массив сэмплов. + """ + sample_rate = ( + self.config.sample_rate + if self.config is not None + else 16000 + ) + if sample_rate <= 0: + sample_rate = 16000 + total_samples = sum(len(chunk) for chunk in self.buffer) + return total_samples / sample_rate + + async def add_audio(self, frame: np.ndarray) -> bool: + """ + Добавляет фрагмент аудио в буфер, если не превышен лимит длительности. + + Args: + frame: Одномерный numpy-массив аудио-сэмплов (float32 или int16). + + Returns: + True — фрагмент успешно добавлен. + False — буфер переполнен, фрагмент отклонён. + """ + self.last_activity = time.time() + + if self.state == SessionState.connecting: + self.state = SessionState.receiving + + # Проверка на переполнение: оцениваем длительность после добавления + sample_rate = ( + self.config.sample_rate + if self.config is not None + else 16000 + ) + if sample_rate <= 0: + sample_rate = 16000 + + incoming_duration = len(frame) / sample_rate + if self.current_buffer_duration_sec + incoming_duration > self.max_buffer_duration_sec: + return False + + self.buffer.append(frame) + return True + + async def get_full_audio(self) -> np.ndarray: + """ + Конкатенирует все фрагменты буфера в единый numpy-массив. + + Returns: + Объединённый массив сэмплов. Если буфер пуст — возвращается пустой массив. + """ + if not self.buffer: + return np.array([], dtype=np.float32) + return np.concatenate(list(self.buffer)) + + async def reset(self) -> None: + """ + Очищает буфер, сбрасывает конфигурацию и переводит сессию + в начальное состояние (connecting). + """ + self.buffer.clear() + self.config = None + self.state = SessionState.connecting + self.last_activity = time.time() + + def is_expired(self, timeout_sec: float) -> bool: + """ + Проверяет, истёк ли таймаут неактивности сессии. + + Args: + timeout_sec: Допустимое время простоя в секундах. + + Returns: + True, если с момента last_activity прошло больше timeout_sec. + """ + return (time.time() - self.last_activity) > timeout_sec diff --git a/tests/test_ws_handler.py b/tests/test_ws_handler.py new file mode 100644 index 0000000..a096476 --- /dev/null +++ b/tests/test_ws_handler.py @@ -0,0 +1,212 @@ +""" +Тесты для services/ws_handler.py (MessageRouter и хендлеры). +""" + +import base64 + +import numpy as np +import pytest + +from models.ws_models import ( + WSConfigMessage, + WSAudioMessage, + WSStatusRequest, + WSPingMessage, + WSPongMessage, + WSErrorMessage, + WSStatusResponse, + WSMessageType, +) +from services.ws_session import AudioSession, SessionState +from services.ws_metrics import SystemMetricsCollector +from services.ws_handler import ( + MessageRouter, + handle_config, + handle_audio, + handle_status_request, + handle_ping, +) + + +class FakeManager: + """ + Фейковый менеджер соединений для тестирования хендлеров. + + Имитирует WSManagerProtocol без реальных WebSocket-объектов. + """ + + def __init__(self, max_connections: int = 100) -> None: + self.sent_messages: list[tuple[str, object]] = [] + self._active_connections = 0 + self._max_connections = max_connections + + async def send_message(self, client_id: str, message: object) -> None: + """Сохраняет сообщение во внутренний список для проверки в тестах.""" + self.sent_messages.append((client_id, message)) + + @property + def active_connections_count(self) -> int: + return self._active_connections + + @property + def max_connections(self) -> int: + return self._max_connections + + +@pytest.fixture +def session() -> AudioSession: + """Фикстура: свежая AudioSession с лимитом буфера 5 секунд.""" + return AudioSession(client_id="test-client", max_buffer_duration_sec=5.0) + + +@pytest.fixture +def manager() -> FakeManager: + """Фикстура: фейковый менеджер соединений.""" + return FakeManager() + + +@pytest.fixture +def metrics() -> SystemMetricsCollector: + """Фикстура: коллектор метрик.""" + return SystemMetricsCollector() + + +class TestHandleConfig: + @pytest.mark.asyncio + async def test_sets_config_and_state(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет, что handle_config устанавливает конфиг и переводит сессию в receiving.""" + msg = WSConfigMessage(sample_rate=8000, audio_transport="binary") + await handle_config(msg, session, manager) + assert session.config == msg + assert session.state == SessionState.receiving + + +class TestHandleAudio: + @pytest.mark.asyncio + async def test_decodes_and_adds_audio(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет успешное декодирование base64 и добавление в буфер.""" + raw = np.zeros(16000, dtype=np.int16).tobytes() + b64 = base64.b64encode(raw).decode() + msg = WSAudioMessage(audio_base64=b64, seq_num=0) + await handle_audio(msg, session, manager) + assert len(session.buffer) == 1 + assert manager.sent_messages == [] + + @pytest.mark.asyncio + async def test_buffer_overflow_sends_error(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет, что при переполнении буфера отправляется WSErrorMessage.""" + raw = np.zeros(96000, dtype=np.int16).tobytes() # 6 сек при 16 кГц, лимит 5 + b64 = base64.b64encode(raw).decode() + msg = WSAudioMessage(audio_base64=b64, seq_num=0) + await handle_audio(msg, session, manager) + assert len(manager.sent_messages) == 1 + client_id, error = manager.sent_messages[0] + assert client_id == session.client_id + assert isinstance(error, WSErrorMessage) + assert error.code == "buffer_overflow" + + @pytest.mark.asyncio + async def test_invalid_base64_sends_error(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет обработку некорректного base64.""" + msg = WSAudioMessage(audio_base64="!!!invalid!!!", seq_num=0) + await handle_audio(msg, session, manager) + assert len(manager.sent_messages) == 1 + _, error = manager.sent_messages[0] + assert isinstance(error, WSErrorMessage) + assert error.code == "decode_error" + + +class TestHandleStatusRequest: + @pytest.mark.asyncio + async def test_sends_status_response( + self, + session: AudioSession, + manager: FakeManager, + metrics: SystemMetricsCollector, + ) -> None: + """Проверяет, что handle_status_request отправляет WSStatusResponse.""" + msg = WSStatusRequest() + await handle_status_request(msg, session, manager, metrics) + assert len(manager.sent_messages) == 1 + _, response = manager.sent_messages[0] + assert isinstance(response, WSStatusResponse) + + +class TestHandlePing: + @pytest.mark.asyncio + async def test_sends_pong(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет, что handle_ping отправляет WSPongMessage.""" + msg = WSPingMessage() + await handle_ping(msg, session, manager) + assert len(manager.sent_messages) == 1 + _, response = manager.sent_messages[0] + assert isinstance(response, WSPongMessage) + + +class TestMessageRouter: + @pytest.mark.asyncio + async def test_register_and_route_config(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет регистрацию хендлера и маршрутизацию сообщения config.""" + router = MessageRouter() + router.register_handler(WSMessageType.config, handle_config) + msg = WSConfigMessage() + await router.route(msg, session, manager) + assert session.config is not None + + @pytest.mark.asyncio + async def test_route_unknown_type_sends_error(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет, что незарегистрированный тип сообщения вызывает ошибку.""" + router = MessageRouter() + msg = WSPingMessage() # ping не зарегистрирован + await router.route(msg, session, manager) + assert len(manager.sent_messages) == 1 + _, error = manager.sent_messages[0] + assert isinstance(error, WSErrorMessage) + assert error.code == "unsupported_type" + + @pytest.mark.asyncio + async def test_route_status_request_without_collector( + self, + session: AudioSession, + manager: FakeManager, + ) -> None: + """Проверяет, что status_request без metrics_collector вызывает handler_error.""" + router = MessageRouter() + router.register_handler(WSMessageType.status_request, handle_status_request) + msg = WSStatusRequest() + await router.route(msg, session, manager, metrics_collector=None) + assert len(manager.sent_messages) == 1 + _, error = manager.sent_messages[0] + assert isinstance(error, WSErrorMessage) + assert error.code == "handler_error" + + @pytest.mark.asyncio + async def test_route_with_collector( + self, + session: AudioSession, + manager: FakeManager, + metrics: SystemMetricsCollector, + ) -> None: + """Проверяет корректную маршрутизацию status_request с metrics_collector.""" + router = MessageRouter() + router.register_handler(WSMessageType.status_request, handle_status_request) + msg = WSStatusRequest() + await router.route(msg, session, manager, metrics_collector=metrics) + assert len(manager.sent_messages) == 1 + _, response = manager.sent_messages[0] + assert isinstance(response, WSStatusResponse) + + @pytest.mark.asyncio + async def test_handler_exception_caught(self, session: AudioSession, manager: FakeManager) -> None: + """Проверяет, что исключение в хендлере перехватывается и отправляется ошибка.""" + async def bad_handler(msg, sess, mgr): + raise ValueError("boom") + + router = MessageRouter() + router.register_handler(WSMessageType.ping, bad_handler) + msg = WSPingMessage() + await router.route(msg, session, manager) + assert len(manager.sent_messages) == 1 + _, error = manager.sent_messages[0] + assert isinstance(error, WSErrorMessage) + assert error.code == "handler_error" diff --git a/tests/test_ws_manager.py b/tests/test_ws_manager.py new file mode 100644 index 0000000..038f335 --- /dev/null +++ b/tests/test_ws_manager.py @@ -0,0 +1,143 @@ +""" +Тесты для services/ws_manager.py (ConnectionManager). +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock + +from services.ws_manager import ConnectionManager, ConnectionMeta +from models.ws_models import WSStatusResponse, WSPingMessage + + +class FakeWebSocket: + """ + Фейковый WebSocket для тестирования ConnectionManager. + Имитирует минимальный интерфейс FastAPI WebSocket. + """ + + def __init__(self): + self.accepted = False + self.closed = False + self.close_code: int | None = None + self.close_reason: str | None = None + self.sent_texts: list[str] = [] + + async def accept(self): + self.accepted = True + + async def close(self, code: int = 1000, reason: str = ""): + self.closed = True + self.close_code = code + self.close_reason = reason + + async def send_text(self, data: str): + self.sent_texts.append(data) + + +@pytest.fixture +def manager() -> ConnectionManager: + """Фикстура: ConnectionManager с лимитом 3 соединения.""" + return ConnectionManager(max_connections=3) + + +class TestConnectionManagerConnect: + @pytest.mark.asyncio + async def test_connect_success(self, manager: ConnectionManager) -> None: + """Успешное подключение добавляет клиента в реестр.""" + ws = FakeWebSocket() + result = await manager.connect(ws, "client-1") + assert result is True + assert ws.accepted is True + assert manager.active_connections_count == 1 + assert "client-1" in manager.active_connections + + @pytest.mark.asyncio + async def test_connect_rejects_when_full(self, manager: ConnectionManager) -> None: + """При превышении лимита соединение отклоняется с кодом 1008.""" + for i in range(3): + ws = FakeWebSocket() + assert await manager.connect(ws, f"client-{i}") is True + + ws_rejected = FakeWebSocket() + result = await manager.connect(ws_rejected, "client-overflow") + assert result is False + assert ws_rejected.closed is True + assert ws_rejected.close_code == 1008 + + +class TestConnectionManagerDisconnect: + @pytest.mark.asyncio + async def test_disconnect_removes_client(self, manager: ConnectionManager) -> None: + """Отключение удаляет клиента из реестров.""" + ws = FakeWebSocket() + await manager.connect(ws, "client-a") + await manager.disconnect("client-a") + assert manager.active_connections_count == 0 + assert "client-a" not in manager.active_connections + assert ws.closed is True + + @pytest.mark.asyncio + async def test_disconnect_unknown_client(self, manager: ConnectionManager) -> None: + """Отключение несуществующего клиента не вызывает ошибок.""" + await manager.disconnect("ghost") + assert manager.active_connections_count == 0 + + +class TestConnectionManagerSendMessage: + @pytest.mark.asyncio + async def test_send_pydantic_model(self, manager: ConnectionManager) -> None: + """Отправка Pydantic-модели сериализуется в JSON.""" + ws = FakeWebSocket() + await manager.connect(ws, "client-b") + msg = WSPingMessage() + await manager.send_message("client-b", msg) + assert len(ws.sent_texts) == 1 + assert '"type":"ping"' in ws.sent_texts[0] + + @pytest.mark.asyncio + async def test_send_dict(self, manager: ConnectionManager) -> None: + """Отправка dict сериализуется в JSON.""" + ws = FakeWebSocket() + await manager.connect(ws, "client-c") + await manager.send_message("client-c", {"foo": "bar"}) + assert len(ws.sent_texts) == 1 + assert '"foo":"bar"' in ws.sent_texts[0] + + @pytest.mark.asyncio + async def test_send_to_unknown_client(self, manager: ConnectionManager) -> None: + """Отправка несуществующему клиенту не вызывает ошибок.""" + await manager.send_message("unknown", WSPingMessage()) + + +class TestConnectionManagerBroadcast: + @pytest.mark.asyncio + async def test_broadcast_status_only_subscribed(self, manager: ConnectionManager) -> None: + """Broadcast отправляет статус только подписанным клиентам.""" + ws1 = FakeWebSocket() + ws2 = FakeWebSocket() + await manager.connect(ws1, "sub-1") + await manager.connect(ws2, "sub-2") + manager.connection_meta["sub-1"].subscribe_status = True + manager.connection_meta["sub-2"].subscribe_status = False + + status = WSStatusResponse(adapter_status="idle") + await manager.broadcast_status(status) + assert len(ws1.sent_texts) == 1 + assert len(ws2.sent_texts) == 0 + + +class TestConnectionManagerDisconnectAll: + @pytest.mark.asyncio + async def test_disconnect_all_closes_everyone(self, manager: ConnectionManager) -> None: + """disconnect_all закрывает все соединения и очищает реестры.""" + ws1 = FakeWebSocket() + ws2 = FakeWebSocket() + await manager.connect(ws1, "all-1") + await manager.connect(ws2, "all-2") + + await manager.disconnect_all(code=1001, reason="shutdown") + assert ws1.closed is True + assert ws2.closed is True + assert ws1.close_code == 1001 + assert manager.active_connections_count == 0 + assert len(manager.connection_meta) == 0 diff --git a/tests/test_ws_metrics.py b/tests/test_ws_metrics.py new file mode 100644 index 0000000..5e58f00 --- /dev/null +++ b/tests/test_ws_metrics.py @@ -0,0 +1,106 @@ +""" +Тесты для services/ws_metrics.py (SystemMetricsCollector). +""" + +import time + +import pytest + +from services.ws_metrics import SystemMetricsCollector +from models.ws_models import WSStatusResponse + + +class TestSystemMetricsCollectorLifecycle: + def test_init_default_start_time(self): + """Проверяет, что start_time устанавливается при инициализации.""" + collector = SystemMetricsCollector() + assert collector.start_time <= time.time() + + def test_init_custom_start_time(self): + """Проверяет передачу кастомного start_time.""" + now = time.time() - 100 + collector = SystemMetricsCollector(start_time=now) + assert collector.start_time == now + + def test_get_gpu_stats_without_handle(self): + """Без NVML-handle GPU-метрики возвращают None.""" + collector = SystemMetricsCollector(nvml_handle=None) + free, total = collector.get_gpu_stats() + assert free is None + assert total is None + + +class TestSystemMetricsCollectorTasks: + def test_active_tasks_counter(self): + """Проверка инкремента/декремента счётчика задач.""" + collector = SystemMetricsCollector() + assert collector.get_active_tasks_count() == 0 + collector.increment_tasks() + assert collector.get_active_tasks_count() == 1 + collector.increment_tasks() + assert collector.get_active_tasks_count() == 2 + collector.decrement_tasks() + assert collector.get_active_tasks_count() == 1 + collector.decrement_tasks() + collector.decrement_tasks() # не уходит ниже 0 + assert collector.get_active_tasks_count() == 0 + + def test_get_queue_depth_default(self): + """По умолчанию глубина очереди равна 0.""" + collector = SystemMetricsCollector() + assert collector.get_queue_depth() == 0 + + +class TestSystemMetricsCollectorStatus: + def test_get_adapter_status_idle(self): + """Статус idle при отсутствии соединений и задач.""" + collector = SystemMetricsCollector() + assert collector.get_adapter_status(0, 100) == "idle" + + def test_get_adapter_status_busy_by_tasks(self): + """Статус busy при наличии активных задач.""" + collector = SystemMetricsCollector() + collector.increment_tasks() + assert collector.get_adapter_status(0, 100) == "busy" + + def test_get_adapter_status_busy_by_connections(self): + """Статус busy при высокой загрузке соединений (порог 80%).""" + collector = SystemMetricsCollector() + assert collector.get_adapter_status(85, 100) == "busy" + + def test_get_adapter_status_idle_low_connections(self): + """Статус idle при низкой загрузке соединений.""" + collector = SystemMetricsCollector() + assert collector.get_adapter_status(50, 100) == "idle" + + +class TestSystemMetricsCollectorCollect: + def test_collect_returns_ws_status_response(self): + """collect() возвращает валидный WSStatusResponse.""" + collector = SystemMetricsCollector(start_time=time.time() - 10) + collector.increment_tasks() + response = collector.collect(active_connections=5, max_connections=100) + assert isinstance(response, WSStatusResponse) + assert response.adapter_status == "busy" + assert response.active_tasks_count == 1 + assert response.active_connections_count == 5 + assert response.uptime_sec >= 10 + assert response.queue_depth == 0 + + def test_collect_zero_connections_idle(self): + """collect() при нулевых соединениях и задачах возвращает idle.""" + collector = SystemMetricsCollector() + response = collector.collect(active_connections=0, max_connections=100) + assert response.adapter_status == "idle" + assert response.active_tasks_count == 0 + assert response.active_connections_count == 0 + assert response.uptime_sec >= 0 + + def test_collect_cpu_stats_optional(self): + """CPU-метрики либо числа, либо None — не вызывают ошибок.""" + collector = SystemMetricsCollector() + response = collector.collect() + if response.cpu_memory_free_mb is not None: + assert isinstance(response.cpu_memory_free_mb, int) + if response.cpu_memory_total_mb is not None: + assert isinstance(response.cpu_memory_total_mb, int) diff --git a/tests/test_ws_models.py b/tests/test_ws_models.py new file mode 100644 index 0000000..0aa8a1b --- /dev/null +++ b/tests/test_ws_models.py @@ -0,0 +1,253 @@ +import json +import pytest +from pydantic import ValidationError + +from models.ws_models import ( + WSMessageType, + WSConfigMessage, + WSAudioMessage, + WSStatusRequest, + WSStatusResponse, + WSResultMessage, + WSWordItem, + WSRecognitionData, + WSErrorMessage, + WSEosMessage, + WSPingMessage, + WSPongMessage, + wrap_binary_audio, + ws_message_adapter, + parse_ws_message, +) + + +class TestWSConfigMessage: + def test_defaults(self): + msg = WSConfigMessage() + assert msg.type == WSMessageType.config + assert msg.sample_rate == 16000 + assert msg.wait_null_answers is True + assert msg.enable_diarization is False + assert msg.audio_transport == "json_base64" + + def test_audio_transport_binary(self): + raw = '{"type": "config", "audio_transport": "binary"}' + msg = WSConfigMessage.model_validate_json(raw) + assert msg.audio_transport == "binary" + + def test_from_json(self): + raw = '{"type": "config", "sample_rate": 8000, "enable_punctuation": true}' + msg = WSConfigMessage.model_validate_json(raw) + assert msg.sample_rate == 8000 + assert msg.enable_punctuation is True + + def test_invalid_sample_rate_too_high(self): + with pytest.raises(ValidationError): + WSConfigMessage(sample_rate=96000) + + def test_invalid_sample_rate_too_low(self): + with pytest.raises(ValidationError): + WSConfigMessage(sample_rate=4000) + + +class TestWSAudioMessage: + def test_from_json(self): + raw = '{"type": "audio_chunk", "audio_base64": "YWJj", "seq_num": 5}' + msg = WSAudioMessage.model_validate_json(raw) + assert msg.seq_num == 5 + assert msg.audio_base64 == "YWJj" + + def test_seq_num_must_be_non_negative(self): + with pytest.raises(ValidationError): + WSAudioMessage(seq_num=-1) + + +class TestWSStatusRequest: + def test_from_json(self): + raw = '{"type": "status_request"}' + msg = WSStatusRequest.model_validate_json(raw) + assert msg.command == "get_status" + + +class TestWSStatusResponse: + def test_from_json(self): + raw = ( + '{"type": "status_response", "adapter_status": "busy", ' + '"gpu_memory_free_mb": 1024, "active_tasks_count": 2}' + ) + msg = WSStatusResponse.model_validate_json(raw) + assert msg.adapter_status == "busy" + assert msg.gpu_memory_free_mb == 1024 + assert msg.active_tasks_count == 2 + + def test_invalid_adapter_status(self): + with pytest.raises(ValidationError): + WSStatusResponse(adapter_status="dead") + + def test_serialization(self): + msg = WSStatusResponse( + adapter_status="overloaded", + gpu_memory_free_mb=512, + gpu_memory_total_mb=4096, + active_connections_count=50, + uptime_sec=123.45, + ) + payload = json.loads(msg.model_dump_json()) + assert payload["adapter_status"] == "overloaded" + assert payload["gpu_memory_free_mb"] == 512 + assert payload["active_connections_count"] == 50 + assert "timestamp" in payload + + +class TestWSWordItem: + def test_from_dict(self): + item = WSWordItem(conf=1.0, start=80.05, end=80.09, word="до") + assert item.word == "до" + assert item.conf == 1.0 + + +class TestWSRecognitionData: + def test_defaults(self): + data = WSRecognitionData() + assert data.result == [] + assert data.text == "" + + def test_with_words(self): + data = WSRecognitionData( + result=[WSWordItem(conf=1.0, start=0.0, end=1.0, word="привет")], + text="привет", + ) + assert len(data.result) == 1 + assert data.text == "привет" + + +class TestWSResultMessage: + def test_partial_result_defaults(self): + data = WSRecognitionData( + result=[WSWordItem(conf=1.0, start=0.0, end=1.0, word="привет")], + text="привет", + ) + msg = WSResultMessage(data=data, last_message=False) + assert msg.silence is False + assert msg.type == WSMessageType.partial_result + assert msg.channel_name == "Null" + assert msg.data.text == "привет" + + def test_final_result(self): + data = WSRecognitionData( + result=[ + WSWordItem(conf=1.0, start=80.05, end=80.09, word="до"), + WSWordItem(conf=1.0, start=80.21, end=80.57, word="свидания"), + ], + text="до свидания", + ) + msg = WSResultMessage( + channel_name="channel_1", + silence=False, + data=data, + last_message=True, + sentenced_data={}, + ) + assert msg.last_message is True + assert msg.data.text == "до свидания" + assert len(msg.data.result) == 2 + + def test_from_json(self): + raw = ( + '{"type": "final_result", "channel_name": "ch1", "silence": false, ' + '"data": {"result": [{"conf": 1, "start": 1.0, "end": 2.0, "word": "test"}], ' + '"text": "test"}, "error": null, "last_message": true, "sentenced_data": {}}' + ) + msg = WSResultMessage.model_validate_json(raw) + assert msg.channel_name == "ch1" + assert msg.data.text == "test" + assert msg.data.result[0].word == "test" + + +class TestWSErrorMessage: + def test_parse_error(self): + msg = WSErrorMessage(code="parse_error", message="Invalid JSON", is_fatal=False) + assert msg.code == "parse_error" + assert not msg.is_fatal + + def test_fatal_error(self): + msg = WSErrorMessage(code="internal_error", message="GPU OOM", is_fatal=True) + assert msg.is_fatal is True + + +class TestWSEosMessage: + def test_type(self): + msg = WSEosMessage() + assert msg.type == WSMessageType.eos + + +class TestWrapBinaryAudio: + def test_wrap_bytes(self): + raw = b"\x00\x01\x02" + msg = wrap_binary_audio(raw, seq_num=7) + assert msg.type == WSMessageType.audio_chunk + assert msg.seq_num == 7 + assert msg.audio_base64 == "AAEC" + + def test_wrap_empty(self): + msg = wrap_binary_audio(b"", seq_num=0) + assert msg.audio_base64 == "" + + +class TestWSMessageDiscriminatedUnion: + def test_union_config(self): + raw = '{"type": "config", "sample_rate": 16000}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSConfigMessage) + + def test_union_audio(self): + raw = '{"type": "audio_chunk", "seq_num": 0}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSAudioMessage) + + def test_union_ping(self): + raw = '{"type": "ping"}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSPingMessage) + + def test_union_pong(self): + raw = '{"type": "pong"}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSPongMessage) + + def test_union_status_request(self): + raw = '{"type": "status_request"}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSStatusRequest) + + def test_union_error(self): + raw = '{"type": "error", "code": "x", "message": "m"}' + msg = ws_message_adapter.validate_json(raw) + assert isinstance(msg, WSErrorMessage) + + def test_union_unknown_type_raises(self): + raw = '{"type": "foobar"}' + with pytest.raises(ValidationError): + ws_message_adapter.validate_json(raw) + + def test_union_invalid_json_raises(self): + raw = "not a json" + with pytest.raises(ValidationError): + ws_message_adapter.validate_json(raw) + + +class TestParseWsMessageHelper: + def test_parse_from_str(self): + raw = '{"type": "pong"}' + msg = parse_ws_message(raw) + assert isinstance(msg, WSPongMessage) + + def test_parse_from_dict(self): + raw = {"type": "eos"} + msg = parse_ws_message(raw) + assert isinstance(msg, WSEosMessage) + + def test_parse_from_bytes(self): + raw = b'{"type": "status_request"}' + msg = parse_ws_message(raw) + assert isinstance(msg, WSStatusRequest) diff --git a/tests/test_ws_session.py b/tests/test_ws_session.py new file mode 100644 index 0000000..bcd43a9 --- /dev/null +++ b/tests/test_ws_session.py @@ -0,0 +1,114 @@ +""" +Тесты для services/ws_session.py (AudioSession). +""" + +import time + +import numpy as np +import pytest + +from services.ws_session import AudioSession, SessionState +from models.ws_models import WSConfigMessage + + +@pytest.fixture +def session(): + """Фикстура: свежая сессия с малым лимитом буфера для тестов.""" + return AudioSession(client_id="test-client-001", max_buffer_duration_sec=2.0) + + +class TestAudioSessionLifecycle: + def test_initial_state(self, session): + """Проверяет начальное состояние сессии после создания.""" + assert session.state == SessionState.connecting + assert session.config is None + assert session.client_id == "test-client-001" + assert len(session.buffer) == 0 + + @pytest.mark.asyncio + async def test_add_audio_success(self, session): + """Успешное добавление аудио и переход в состояние receiving.""" + # 1 секунда аудио при 16 кГц = 16000 сэмплов + frame = np.zeros(16000, dtype=np.float32) + result = await session.add_audio(frame) + assert result is True + assert session.state == SessionState.receiving + assert len(session.buffer) == 1 + + @pytest.mark.asyncio + async def test_add_audio_overflow(self, session): + """Отказ при превышении max_buffer_duration_sec.""" + # Лимит 2 секунды. Первый фрейм 1.2 сек — ок. Второй 1.2 сек — переполнение. + frame = np.zeros(int(1.2 * 16000), dtype=np.float32) + assert await session.add_audio(frame) is True + assert await session.add_audio(frame) is False + assert len(session.buffer) == 1 + + @pytest.mark.asyncio + async def test_add_audio_with_config(self, session): + """Проверка учёта sample_rate из конфига при расчёте длительности.""" + session.config = WSConfigMessage(sample_rate=8000) + frame = np.zeros(8000, dtype=np.float32) # 1 сек при 8 кГц + assert await session.add_audio(frame) is True + assert session.current_buffer_duration_sec == pytest.approx(1.0, 0.01) + + @pytest.mark.asyncio + async def test_get_full_audio(self, session): + """Конкатенация нескольких фреймов в единый массив.""" + frame1 = np.ones(8000, dtype=np.float32) + frame2 = np.ones(8000, dtype=np.float32) * 2 + await session.add_audio(frame1) + await session.add_audio(frame2) + full = await session.get_full_audio() + assert len(full) == 16000 + assert full[0] == 1.0 + assert full[-1] == 2.0 + + @pytest.mark.asyncio + async def test_get_full_audio_empty_buffer(self, session): + """Пустой буфер возвращает пустой numpy-массив.""" + full = await session.get_full_audio() + assert len(full) == 0 + assert full.dtype == np.float32 + + @pytest.mark.asyncio + async def test_reset(self, session): + """Сброс сессии в начальное состояние.""" + session.config = WSConfigMessage(sample_rate=16000) + await session.add_audio(np.zeros(16000, dtype=np.float32)) + session.state = SessionState.processing + + await session.reset() + assert session.state == SessionState.connecting + assert session.config is None + assert len(session.buffer) == 0 + assert session.current_buffer_duration_sec == 0.0 + + def test_is_expired_true(self, session): + """Таймаут истёк.""" + session.last_activity = time.time() - 100 + assert session.is_expired(timeout_sec=10) is True + + def test_is_expired_false(self, session): + """Таймаут не истёк.""" + session.last_activity = time.time() + assert session.is_expired(timeout_sec=10) is False + + +class TestAudioSessionBufferDuration: + @pytest.mark.asyncio + async def test_duration_calculation_with_different_sample_rates(self): + """Суммарная длительность при стандартном sample_rate.""" + sess = AudioSession(client_id="sr-test", max_buffer_duration_sec=5.0) + sess.config = WSConfigMessage(sample_rate=16000) + await sess.add_audio(np.zeros(8000, dtype=np.float32)) # 0.5 сек + await sess.add_audio(np.zeros(8000, dtype=np.float32)) # ещё 0.5 сек + assert sess.current_buffer_duration_sec == pytest.approx(1.0, 0.01) + + @pytest.mark.asyncio + async def test_duration_respects_config_sample_rate(self): + """Суммарная длительность пересчитывается при sample_rate=8000.""" + sess = AudioSession(client_id="sr-test-8k", max_buffer_duration_sec=5.0) + sess.config = WSConfigMessage(sample_rate=8000) + await sess.add_audio(np.zeros(8000, dtype=np.float32)) # 1 сек при 8кГц + assert sess.current_buffer_duration_sec == pytest.approx(1.0, 0.01) From 7b6d2a2d9bd8c92258de8909120f43d3e375bacb Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 7 May 2026 12:21:09 +0300 Subject: [PATCH 27/46] =?UTF-8?q?=D0=92=D1=8B=D0=B2=D0=BE=D0=B4=D0=B8?= =?UTF-8?q?=D0=BC=20gl=D0=BEbals=20=D0=B8=D0=B7=20url=20=D0=B8=20file.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/engine/file_recognition.py | 163 +++++------- api/v1/endpoints/asr_ws.py | 367 +++++++++++--------------- config.py | 183 ++++++++----- core/state_store.py | 97 +++++++ main.py | 13 + models/ws_models.py | 2 + routes/ws_audio_transkrib.py | 3 +- services/asr_pipeline.py | 264 ++++++++++++++++++ services/recognition_session.py | 161 +++++++++++ services/ws_session.py | 70 +++++ templates/index.html | 44 ++- tests/test_asr_pipeline.py | 233 ++++++++++++++++ tests/test_chunk_doing_v2.py | 47 ++++ tests/test_recognition_session.py | 96 +++++++ tests/test_state_store.py | 92 +++++++ tests/test_ws_models.py | 6 + tests/test_ws_session.py | 27 ++ 17 files changed, 1469 insertions(+), 399 deletions(-) create mode 100644 core/state_store.py create mode 100644 services/asr_pipeline.py create mode 100644 services/recognition_session.py create mode 100644 tests/test_asr_pipeline.py create mode 100644 tests/test_chunk_doing_v2.py create mode 100644 tests/test_recognition_session.py create mode 100644 tests/test_state_store.py diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index ce3cc7f..a31ea33 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -1,32 +1,27 @@ -from fastapi import Depends import time from pydub import AudioSegment from config import settings import asyncio -import uuid import logging -from utils.pre_start_init import ( - posted_and_downloaded_audio, - audio_buffer, - audio_overlap, - audio_to_asr, - audio_duration, -) - -from utils.chunk_doing import find_last_speech_position + +from services.recognition_session import FileRecognitionSession +from utils.chunk_doing import find_last_speech_position_v2 from utils.resamppling import sync_resample_audiosegment from Recognizer.engine.stream_recognition import simple_recognise, recognise_w_speed_correction, simple_recognise_batch from Recognizer.engine.sentensizer import do_sensitizing from Recognizer.engine.echoe_clearing import remove_echo from Diarisation.diarazer import do_diarizing -from threading import Lock logger = logging.getLogger(__name__) -# Глобальный лок для потокобезопасности -audio_lock = Lock() +def process_file(session=None, recognizer=None, punctuator=None, diarizer=None, tmp_path=None, params=None): + # Legacy adapter: поддержка старых роутов, передающих tmp_path + params + if session is None and tmp_path is not None and params is not None: + session = FileRecognitionSession(params=params) + session.tmp_path = tmp_path + elif session is None: + raise TypeError("process_file() требует либо 'session', либо оба 'tmp_path' и 'params'") -def process_file(tmp_path, params, recognizer, punctuator, diarizer): process_file_start = time.perf_counter() res = False diarized = False @@ -39,15 +34,15 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): "sentenced_data": dict(), } - post_id = str(uuid.uuid4()) + post_id = session.post_id + params = session.params logger.debug(f'Принят новый "post_file" id = {post_id}') try: - with audio_lock: - if params.make_mono: - posted_and_downloaded_audio[post_id] = AudioSegment.from_file(tmp_path).set_channels(1) - else: - posted_and_downloaded_audio[post_id] = AudioSegment.from_file(tmp_path) + if params.make_mono: + audio_input = AudioSegment.from_file(session.tmp_path).set_channels(1) + else: + audio_input = AudioSegment.from_file(session.tmp_path) except Exception as e: error_description += f"Error loading audio file: {e}" logger.error(error_description) @@ -57,10 +52,9 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): # Проверка длины переданного на распознавание аудио try: - with audio_lock: - if posted_and_downloaded_audio[post_id].duration_seconds < 5: - logger.debug(f"На вход передано аудио короче 5 секунд. Будет дополнено тишиной ещё 5 сек.") - posted_and_downloaded_audio[post_id] += AudioSegment.silent(duration=5, frame_rate=settings.BASE_SAMPLE_RATE) + if audio_input.duration_seconds < 5: + logger.debug(f"На вход передано аудио короче 5 секунд. Будет дополнено тишиной ещё 5 сек.") + audio_input += AudioSegment.silent(duration=5, frame_rate=settings.BASE_SAMPLE_RATE) except Exception as e: error_description += f"Error len_fixing_file: {e}" logger.error(error_description) @@ -71,19 +65,12 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): # Приводим фреймрейт к фреймрейту модели logger.debug(f"Начало проверки фреймрейта {(time.perf_counter()-process_file_start):.4f} сек.") try: - with audio_lock: - if posted_and_downloaded_audio[post_id].frame_rate != settings.BASE_SAMPLE_RATE: - posted_and_downloaded_audio[post_id] = sync_resample_audiosegment( - audio_data=posted_and_downloaded_audio[post_id], - target_sample_rate=settings.BASE_SAMPLE_RATE) - logger.debug(f"Корректировка фреймрейта {(time.perf_counter() - process_file_start):.4f} сек.") - except KeyError as e_key: - error_description = f"Ошибка обращения по ключу {post_id} при изменения фреймрейта - {e_key}" - logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) - return result - + if audio_input.frame_rate != settings.BASE_SAMPLE_RATE: + audio_input = sync_resample_audiosegment( + audio_data=audio_input, + target_sample_rate=settings.BASE_SAMPLE_RATE + ) + logger.debug(f"Корректировка фреймрейта {(time.perf_counter() - process_file_start):.4f} сек.") except Exception as e: error_description = f"Ошибка изменения фреймрейта - {e}" logger.error(error_description) @@ -92,61 +79,51 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): return result # Обрабатываем чанки с аудио по N секунд - for n_channel, mono_data in enumerate(posted_and_downloaded_audio[post_id].split_to_mono()): + session.collected_asr_res = {} + for n_channel, mono_data in enumerate(audio_input.split_to_mono()): time_chunks_start = time.perf_counter() # Подготовительные действия - try: - with audio_lock: - audio_buffer[post_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - audio_overlap[post_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - audio_duration[post_id] = 0 - except Exception as e: - error_description = f"Ошибка изменения фреймрейта - {e}" - logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) - return result + session.audio_buffer = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + session.audio_overlap = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + session.audio_to_asr = [] + session.audio_duration = 0.0 - result["raw_data"].update({f"channel_{n_channel + 1}": list()}) + session.collected_asr_res[f"channel_{n_channel + 1}"] = [] # Основной процесс перебора чанков для распознавания overlaps = list(mono_data[::settings.MAX_OVERLAP_DURATION * 1000]) # Чанки аудио для распознавания total_chunks = len(overlaps) # Количество чанков, для поиска последнего - audio_to_asr[post_id] = list() for idx, overlap in enumerate(overlaps): - is_last_chunk = (idx == total_chunks - 1) # Если чанк последний - with audio_lock: - if (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) < settings.MAX_OVERLAP_DURATION: - silent_secs = settings.MAX_OVERLAP_DURATION - (audio_overlap[post_id].duration_seconds + overlap.duration_seconds) - overlap += AudioSegment.silent(silent_secs, frame_rate=settings.BASE_SAMPLE_RATE) - audio_buffer[post_id] = overlap - asyncio.run(find_last_speech_position(post_id, is_last_chunk)) # Последний чанк обрабатывается иначе. + is_last_chunk = (idx == total_chunks - 1) # Если чанк последний + if (session.audio_overlap.duration_seconds + overlap.duration_seconds) < settings.MAX_OVERLAP_DURATION: + silent_secs = settings.MAX_OVERLAP_DURATION - (session.audio_overlap.duration_seconds + overlap.duration_seconds) + overlap += AudioSegment.silent(silent_secs, frame_rate=settings.BASE_SAMPLE_RATE) + session.audio_buffer = overlap + asyncio.run(find_last_speech_position_v2(session, is_last_chunk)) # Последний чанк обрабатывается иначе. if params.use_batch: logger.info("Запрошен батчинг") - list_asr_result_wo_conf = simple_recognise_batch(audio_to_asr[post_id],params.batch_size, recognizer) # --> list + list_asr_result_wo_conf = simple_recognise_batch(session.audio_to_asr, params.batch_size, recognizer) # --> list - for _, asr_result_wo_conf in enumerate(list_asr_result_wo_conf): + for idx, asr_result_wo_conf in enumerate(list_asr_result_wo_conf): + asr_result = recognizer.apply_postprocessing(asr_result_wo_conf, session.audio_duration) - asr_result = recognizer.apply_postprocessing(asr_result_wo_conf, audio_duration[post_id]) - - result["raw_data"][f"channel_{n_channel + 1}"].append(asr_result) - with audio_lock: - audio_duration[post_id] += audio_to_asr[post_id][_].duration_seconds + session.collected_asr_res[f"channel_{n_channel + 1}"].append(asr_result) + session.audio_duration += session.audio_to_asr[idx].duration_seconds res = True logger.debug(asr_result) else: - for audio_asr in audio_to_asr[post_id]: + for audio_asr in session.audio_to_asr: try: # Снижаем скорость аудио по необходимости if params.do_auto_speech_speed_correction or params.speech_speed_correction_multiplier != 1: logger.debug("Будут использованы механизмы анализа скорости речи и замедления аудио") asr_result_wo_conf, speed, multiplier = asyncio.run(recognise_w_speed_correction( - audio_data=audio_asr, - can_slow_down=True, - multiplier=params.speech_speed_correction_multiplier, - recognizer=recognizer) - ) + audio_data=audio_asr, + can_slow_down=True, + multiplier=params.speech_speed_correction_multiplier, + recognizer=recognizer) + ) params.speech_speed_correction_multiplier = multiplier else: # Производим распознавание @@ -156,32 +133,19 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): logger.error(f"Error ASR audio - {e}") error_description = f"Error ASR audio - {e}" else: - asr_result = recognizer.apply_postprocessing(asr_result_wo_conf, audio_duration[post_id]) + asr_result = recognizer.apply_postprocessing(asr_result_wo_conf, session.audio_duration) - result["raw_data"][f"channel_{n_channel + 1}"].append(asr_result) + session.collected_asr_res[f"channel_{n_channel + 1}"].append(asr_result) - with audio_lock: - audio_duration[post_id] += audio_asr.duration_seconds + session.audio_duration += audio_asr.duration_seconds res = True logger.debug(asr_result) - - with audio_lock: - try: - del audio_overlap[post_id] - del audio_buffer[post_id] - del audio_to_asr[post_id] - del audio_duration[post_id] - except Exception as e: - error_description = f"Ошибка при очистке данных - {e}" - logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) del mono_data del overlaps if params.do_echo_clearing: try: - result["raw_data"] = asyncio.run(remove_echo(result["raw_data"])) + session.collected_asr_res = asyncio.run(remove_echo(session.collected_asr_res)) except Exception as e: logger.error(f"Error echo clearing - {e}") @@ -193,7 +157,7 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): logger.error("Запрошена диаризация, но она не доступна.") params.do_diarization = False # Проверяем возможность диаризации. Если здесь стерео-канал, то диаризацию выключаем. - elif params.do_diarization and len(posted_and_downloaded_audio[post_id].split_to_mono()) != 1: + elif params.do_diarization and len(audio_input.split_to_mono()) != 1: error_description += f"Only mono diarization available.\n" logger.warn("При запрошенной диаризации аудио имеет более одного аудио-канала. Диаризация будет выключена.") params.do_diarization = False @@ -202,7 +166,7 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): try: result["diarized_data"] = asyncio.run(do_diarizing( file_id=str(post_id), - asr_raw_data=result["raw_data"], + asr_raw_data=session.collected_asr_res, diar_vad_sensity=params.diar_vad_sensity, diarizer=diarizer, )) @@ -214,29 +178,28 @@ def process_file(tmp_path, params, recognizer, punctuator, diarizer): diarized = True if params.do_dialogue: - data_to_do_sensitizing = result["diarized_data"] if diarized else result["raw_data"] + data_to_do_sensitizing = result["diarized_data"] if diarized else session.collected_asr_res try: result["sentenced_data"] = asyncio.run(do_sensitizing( - input_asr_json=data_to_do_sensitizing, - do_punctuation=params.do_punctuation, - punctuator=punctuator - ) - ) + input_asr_json=data_to_do_sensitizing, + do_punctuation=params.do_punctuation, + punctuator=punctuator + )) except Exception as e: logger.error(f"do_sensitizing - {e}") error_description = f"do_sensitizing - {e}" res = False else: if not params.keep_raw: - result["raw_data"].clear() + session.collected_asr_res.clear() else: result["sentenced_data"].clear() result["error_description"] = error_description result["success"] = res - with audio_lock: - del posted_and_downloaded_audio[post_id] + # Переносим raw_data из сессии в результат для обратной совместимости ответа + result["raw_data"] = session.collected_asr_res logger.debug(result) return result diff --git a/api/v1/endpoints/asr_ws.py b/api/v1/endpoints/asr_ws.py index c64b765..72abc9d 100644 --- a/api/v1/endpoints/asr_ws.py +++ b/api/v1/endpoints/asr_ws.py @@ -1,244 +1,175 @@ -from pydub import AudioSegment - -import ujson +""" +WebSocket-роут /api/v1/asr/ws +Использует ConnectionManager, AudioSession, MessageRouter, asr_pipeline. +Сохраняет обратную совместимость протокола (config, audio, eof/eos). +""" + +import asyncio +import base64 import logging -from config import settings import uuid -from io import BytesIO +from contextlib import asynccontextmanager -from fastapi import APIRouter, WebSocket, Depends -from utils.chunk_doing import find_last_speech_position -from utils.pre_start_init import audio_buffer, audio_overlap, audio_to_asr, audio_duration, ws_collected_asr_res -from utils.send_messages import send_messages -from utils.tokens_to_Result import process_single_token_vocab_output -from utils.resamppling import async_resample_audiosegment +from fastapi import APIRouter, WebSocket, WebSocketDisconnect, Depends +from config import settings +from models.ws_models import ( + WSConfigMessage, + WSAudioMessage, + WSEosMessage, + WSErrorMessage, + WSMessageType, + parse_ws_message, +) +from services.ws_manager import ConnectionManager +from services.ws_session import AudioSession, SessionState +from services.ws_handler import MessageRouter, handle_config, handle_ping, handle_status_request +from services.ws_metrics import SystemMetricsCollector +from services.asr_pipeline import process_audio_stream_chunk, process_final_audio from Recognizer import get_recognizer, Recognizer -from Recognizer.engine.sentensizer import do_sensitizing -from Recognizer.engine.stream_recognition import simple_recognise from Punctuation import get_punctuator, SbertPuncCaseOnnx router = APIRouter(prefix="/asr", tags=["ASR"]) logger = logging.getLogger(__name__) +@asynccontextmanager +async def audio_session_lifecycle(client_id: str): + """ + Контекстный менеджер жизненного цикла AudioSession. + + Гарантирует очистку AudioSegment-буферов и глобальных dict при выходе. + """ + session = AudioSession(client_id=client_id) + try: + yield session + finally: + await session.reset() + _cleanup_globals(client_id) + + @router.websocket("/ws") -async def websocket( - ws: WebSocket, +async def websocket_endpoint( + websocket: WebSocket, recognizer: Recognizer = Depends(get_recognizer), punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), ): - wait_null_answers = True - client_id = uuid.uuid4() - logger.debug(f'Принят новый сокет id = {client_id}') - audio_buffer[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - audio_overlap[client_id] = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - audio_duration[client_id] = 0 - audio_to_asr[client_id] = list() - ws_collected_asr_res[client_id] = {f"channel_{1}": list()} - do_dialogue = False - do_punctuation = False - audio_format = 'raw' - sample_rate = settings.BASE_SAMPLE_RATE - sentenced_data = None - error_description = None - - await ws.accept() - channel_name = str() - - while True: - try: - message = await ws.receive() - except Exception as wse: - logger.error(f"receive WebSocketException - {wse}") - return - - if isinstance(message, dict) and message.get('text'): - try: - if message.get('text') and 'config' in message.get('text'): - json_cfg = ujson.loads(message.get('text'))['config'] - audio_format = json_cfg.get("audio_format", 'pcm16') - sample_rate = json_cfg.get('sample_rate') - wait_null_answers = json_cfg.get('wait_null_answers', wait_null_answers) - do_dialogue = json_cfg.get("do_dialogue", False) - do_punctuation = json_cfg.get("do_punctuation", False) - try: - channel_name = message.get('text').get("channelName") - except Exception as e: - channel_name = "Null" - logger.debug("ChannelName not parsed") - logger.info(f"Task received, config - {message.get('text')}") - continue + """ + WebSocket endpoint для потокового распознавания речи (ASR). + + Протокол: + 1. Клиент отправляет config (WSConfigMessage). + 2. Клиент отправляет audio_chunk (WSAudioMessage или binary frame). + 3. По завершении — eos/eof (WSEosMessage или текст "eof"). + """ + manager: ConnectionManager = websocket.app.state.ws_manager + metrics: SystemMetricsCollector = websocket.app.state.metrics_collector + state_store = websocket.app.state.state_store + + client_id = str(uuid.uuid4()) + logger.info("New WS connection: %s", client_id) + + # 1. Подключение (с проверкой лимита соединений) + if not await manager.connect(websocket, client_id): + logger.warning("Connection rejected for %s (max connections reached)", client_id) + return + + # 2. Жизненный цикл сессии (гарантированная очистка в finally) + async with audio_session_lifecycle(client_id) as session: + # 3. Регистрация хендлеров сообщений + msg_router = MessageRouter() + msg_router.register_handler(WSMessageType.config, handle_config) + msg_router.register_handler(WSMessageType.ping, handle_ping) + msg_router.register_handler(WSMessageType.status_request, handle_status_request) - elif message.get('text') and 'eof' in message.get('text'): - logger.info(f"EOF received in channel {channel_name}") + try: + while True: + # 4. Получение сообщения с idle timeout + try: + message = await asyncio.wait_for( + websocket.receive(), + timeout=settings.WS_IDLE_TIMEOUT_SEC, + ) + except asyncio.TimeoutError: + logger.info("Idle timeout for client %s", client_id) + await manager.send_message( + client_id, + WSErrorMessage( + code="idle_timeout", + message=f"No messages for {settings.WS_IDLE_TIMEOUT_SEC}s", + is_fatal=False, + ), + ) break - else: - logger.error(f"Can`t recognise text part of message {message.get('text')} in channel {channel_name}") - - except Exception as e: - logger.error(f'Error text message compiling. Message:{message} - error:{e} in channel {channel_name}') - elif isinstance(message, dict) and message.get('bytes'): - try: - chunk = message.get('bytes') - if audio_format == 'pcm16': - if len(chunk) % 2 != 0: - chunk += bytes(2 - (len(chunk) % 2)) + # Обработка disconnect от клиента + if message.get("type") == "websocket.disconnect": + logger.info("Client %s disconnected (code=%s)", client_id, message.get("code")) + break - audiosegment_chunk = AudioSegment( - chunk, - frame_rate=sample_rate, - sample_width=2, - channels=1 + # Определяем тип содержимого: bytes (binary) или text (JSON) + if message.get("bytes"): + # Binary frame: отправляем сырые байты напрямую в pipeline, без base64-обёртки + await process_audio_stream_chunk( + session, message["bytes"], recognizer, punctuator, manager ) - - else: - try: - buffer = BytesIO(chunk) - buffer.seek(0) - audiosegment_chunk = AudioSegment.from_file(buffer) - - except Exception as e: - logger.error(f"Ошибка принятия аудио - {e} in channel {channel_name}") - else: - logger.debug(f"Чанк принят и распознан in channel {channel_name}") - - if audiosegment_chunk.frame_rate != settings.BASE_SAMPLE_RATE: - audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, settings.BASE_SAMPLE_RATE) - - if audiosegment_chunk.channels != 1: - audiosegment_chunk = audiosegment_chunk.set_channels(1) - audio_buffer[client_id] += audiosegment_chunk - - if (audio_overlap[client_id] + audio_buffer[client_id]).duration_seconds >= settings.MAX_OVERLAP_DURATION: - await find_last_speech_position(client_id, is_last_chunk=False) - + continue + elif message.get("text"): + msg = parse_ws_message(message["text"]) else: + logger.warning("Unknown WS message format for %s: %s", client_id, message) continue - except Exception as e: - logger.error(f"AcceptWaveform error - {e} in channel {channel_name}") - else: - try: - asr_result = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) - asr_result_words = process_single_token_vocab_output(asr_result, audio_duration[client_id]) - audio_duration[client_id] += audio_to_asr[client_id][-1].duration_seconds - logger.debug(asr_result_words) - - ws_collected_asr_res[client_id][f"channel_{1}"].append(asr_result_words) - except Exception as e: - logger.error(f"recognizer.get_result(stream()) error - {e}") - else: - if len(asr_result_words.get("data").get("text")) == 0 or asr_result_words.get("data").get("text") == ' ': - if wait_null_answers: - if not await send_messages(ws, _silence=True, _data=None, _error=None, _channel_name=channel_name): - logger.error(f"send_message not ok work canceled") - try: - del audio_overlap[client_id] - del audio_buffer[client_id] - del audio_to_asr[client_id] - del audio_duration[client_id] - del ws_collected_asr_res[client_id] - except Exception as e: - logger.error(f"error clearing globals after abnormal closing socket - {e}") - return - else: - logger.debug("sending silence partials skipped") - continue - else: - if not await send_messages(ws, _silence=False, _data=asr_result_words, _error=None, _channel_name=channel_name): - logger.error(f"send_message not ok work canceled") - try: - del audio_overlap[client_id] - del audio_buffer[client_id] - del audio_to_asr[client_id] - del audio_duration[client_id] - del ws_collected_asr_res[client_id] - except Exception as e: - logger.error(f"error clearing globals after abnormal closing socket - {e}") - return - elif isinstance(message, dict) and message.get('type') == "websocket.disconnect": - description = f"Channel {channel_name} closed from outside" - logger.error(description) - break - else: - error_description = f"Can`t parse message - {message} in channel {channel_name}" - logger.error(error_description) - - if not await send_messages(ws, _silence=False, _data=None, _error=error_description, _channel_name=channel_name): - logger.error(f"send_message not ok work canceled in channel {channel_name}") - try: - del audio_overlap[client_id] - del audio_buffer[client_id] - del audio_to_asr[client_id] - del audio_duration[client_id] - del ws_collected_asr_res[client_id] - except Exception as e: - logger.error(f"error clearing globals after abnormal closing socket - {e} in channel {channel_name}") - return - - audio_to_asr[client_id].append(audio_overlap[client_id] + audio_buffer[client_id]) - logger.debug(f'итоговое сообщение - {audio_to_asr[client_id][-1].duration_seconds} секунд') + # 5. Маршрутизация служебных сообщений + if msg.type in ( + WSMessageType.config, + WSMessageType.ping, + WSMessageType.status_request, + ): + await msg_router.route(msg, session, manager, metrics_collector=metrics) + + # Копирование флагов из конфига в сессию (для ASR pipeline) + if isinstance(msg, WSConfigMessage): + session.wait_null_answers = msg.wait_null_answers + session.do_dialogue = msg.do_dialogue + session.do_punctuation = msg.do_punctuation + session.channel_name = msg.channel_name or "Null" + + # 6. Обработка аудио-чанка + if isinstance(msg, WSAudioMessage): + chunk_bytes = b"" + if msg.audio_base64: + chunk_bytes = base64.b64decode(msg.audio_base64) + if chunk_bytes: + await process_audio_stream_chunk( + session, chunk_bytes, recognizer, punctuator, manager + ) + + # 7. Обработка конца потока (eos/eof) + if isinstance(msg, WSEosMessage): + await process_final_audio(session, recognizer, punctuator, manager) + break - try: - try: - if audio_to_asr[client_id][-1].duration_seconds < 2: - audio_to_asr[client_id][-1] = audio_to_asr[client_id][-1] + AudioSegment.silent(1000, frame_rate=sample_rate) - except Exception as e: - logger.error(f"Ошибка дополнения тишиной последнего чанка - {e} in channel {channel_name}") - last_result = None - error_description = f"Ошибка дополнения тишиной последнего чанка - {e} in channel {channel_name}" - else: - last_asr_result_w_conf = await simple_recognise(audio_to_asr[client_id][-1], recognizer=recognizer) - last_result = process_single_token_vocab_output(last_asr_result_w_conf, audio_duration[client_id]) - logger.debug(f'Последний результат {last_result.get("data").get("text")} in channel {channel_name}') - - ws_collected_asr_res[client_id][f"channel_{1}"].append(last_result) - - except Exception as e: - logger.error(f"last_asr_result_w_conf error - {e}") - - else: - if len(last_result.get("data").get("text")) == 0: - is_silence = True - last_result = None - elif last_result.get("data").get("text") == ' ': - is_silence = True - last_result = None - else: - logger.debug(last_result) - is_silence = False - - if do_dialogue: + except WebSocketDisconnect: + logger.info("Client %s disconnected normally", client_id) + except Exception as exc: + logger.exception("WS error for %s: %s", client_id, exc) try: - sentenced_data = await do_sensitizing(ws_collected_asr_res[client_id], do_punctuation, punctuator=punctuator) - except Exception as e: - logger.error(f"await do_sensitizing - {e}") - error_description = f"do_sensitizing - {e}" - - if not await send_messages(ws, _silence=is_silence, _data=last_result, _error=error_description, _last_message=True, - _sentenced_data=sentenced_data, _channel_name=channel_name): - logger.error(f"send_message not ok work canceled in channel {channel_name}") + await manager.send_message( + client_id, + WSErrorMessage( + code="internal_error", + message=str(exc), + is_fatal=True, + ), + ) + except Exception: + pass + finally: + logger.info("Closing WS connection %s", client_id) + # Сохранение мета-информации в StateStore (аудит / восстановление) try: - del audio_overlap[client_id] - del audio_buffer[client_id] - del audio_to_asr[client_id] - del audio_duration[client_id] - del ws_collected_asr_res[client_id] - except Exception as e: - logger.error(f"error clearing globals after abnormal closing socket - {e} in channel {channel_name}") - return - - logger.info(f"Closing connection {channel_name}") - await ws.close() - - try: - del audio_overlap[client_id] - del audio_buffer[client_id] - del audio_to_asr[client_id] - del audio_duration[client_id] - del ws_collected_asr_res[client_id] - except Exception as e: - logger.error(f"error clearing globals after NORMAL closing socket - {e} in channel {channel_name}") - return + await state_store.set(f"session:{client_id}", session.to_dict()) + except Exception as exc: + logger.debug("Failed to save session state: %s", exc) + await manager.disconnect(client_id) diff --git a/config.py b/config.py index 6495de4..7b96c75 100644 --- a/config.py +++ b/config.py @@ -196,109 +196,148 @@ def _compute_derived(self): AUDIOEXTENTIONS = [ # Основные форматы - '.mp3', '.wav', '.aac', '.ogg', '.flac', '.m4a', '.wma', '.aiff', '.alac', + 'mp3', 'wav', 'aac', 'ogg', 'flac', 'm4a', 'wma', 'aiff', 'alac', # Менее распространённые форматы - '.ape', '.opus', '.amr', '.au', '.mid', '.midi', '.ac3', '.dts', '.ra', '.rm', '.voc', + 'ape', 'opus', 'amr', 'au', 'mid', 'midi', 'ac3', 'dts', 'ra', 'rm', 'voc', # Форматы для сжатия и профессионального аудио - '.dsd', '.pcm', '.raw', '.tta', '.webm', '.3ga', '.8svx', '.cda', + 'dsd', 'pcm', 'raw', 'tta', 'webm', '3ga', '8svx', 'cda', # Форматы с потерями и без потерь - '.mp2', '.mp1', '.gsm', '.vox', '.dss', '.mka', '.tak', '.ofr', '.spx', + 'mp2', 'mp1', 'gsm', 'vox', 'dss', 'mka', 'tak', 'ofr', 'spx', # Игровые аудиоформаты - '.xm', '.mod', '.s3m', '.it', '.nsf', + 'xm', 'mod', 's3m', 'it', 'nsf', # Редкие/устаревшие форматы - '.669', '.mtm', '.med', '.far', '.umx' + '669', 'mtm', 'med', 'far', 'umx' ] # Описание WebSocket для OpenAPI WS_DESCRIPTION = """ -## WebSocket Endpoint - `/ws` +## WebSocket Endpoint — `/api/v1/asr/ws` (актуальный протокол) + ### Пример конфигурации Отправьте JSON с конфигурацией: ```json - { - "config": { - "audio_format": "pcm16", - "sample_rate": 16000, - "wait_null_answers": true, - "do_dialogue": false, - "do_punctuation": false, - "channelName": "channel_1" # id канала из астериск, например. - } - } +{ + "type": "config", + "sample_rate": 16000, + "audio_format": "pcm16", + "audio_transport": "json_base64", + "wait_null_answers": true, + "do_dialogue": false, + "do_punctuation": false, + "channel_name": "channel_1" +} ``` -### Пример передачи данных (JSON + base64) - -Используется при `audio_transport: "json_base64"` (по умолчанию): +### Пример передачи аудио (JSON + base64) ```json - { - "type": "audio_chunk", - "audio_base64": "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAEAfAAABAAgAZGF0YQAAAAA=", - "seq_num": 0 - } +{ + "type": "audio_chunk", + "audio_base64": "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAEAfAAABAAgAZGF0YQAAAAA=", + "seq_num": 0 +} ``` -### Пример передачи данных (Binary) +### Пример передачи аудио (Binary) Используется при `audio_transport: "binary"`. Отправляйте WebSocket **binary frame** напрямую (без JSON-обёртки). Сервер читает его через `receive_bytes()`. -### Пример EOF +### Пример завершения потока (EOS) + +```json +{ + "type": "eos" +} +``` + +### Формат ответа + +```json +{ + "type": "partial_result", + "channel_name": "channel_1", + "silence": false, + "data": { + "result": [ + {"conf": 1.0, "start": 116.48, "end": 116.76, "word": "владимир"}, + {"conf": 1.0, "start": 116.92, "end": 117.48, "word": "анатольевич"} + ], + "text": "владимир анатольевич" + }, + "error": null, + "last_message": false, + "sentenced_data": null +} +``` + +При завершении (`last_message: true`) и включённых `do_dialogue` + `do_punctuation`: -По завершении отправьте: ```json - { - "text": "eof" +{ + "type": "final_result", + "channel_name": "channel_1", + "silence": false, + "data": { + "result": [...], + "text": "..." + }, + "error": null, + "last_message": true, + "sentenced_data": { + "raw_text_sentenced_recognition": "channel_1: Ничьих, не требуя ... мои.\\nchannel_1: У Лукоморья дуб зеленый.", + "list_of_sentenced_recognitions": [...], + "full_text_only": ["..."] } +} ``` -### Ответы - -``` json - { - 'channel_name': 'Null', - 'silence': False, - 'data': { - 'result': - [ - {'conf': 1.0, 'start': 116.48, 'end': 116.76, 'word': 'владимир'}, - {'conf': 1.0, 'start': 116.92, 'end': 117.48, 'word': 'анатольевич'} - ], - 'text': 'владимир анатольевич'}, - 'error': None, - 'last_message': False, - 'sentenced_data': {} - } +--- + +## WebSocket Endpoint — `/ws` (legacy, deprecated) + +> ⚠️ **Deprecated**: этот endpoint сохраняется для обратной совместимости. Используйте `/api/v1/asr/ws`. + +### Пример конфигурации (legacy) + +```json +{ + "config": { + "audio_format": "pcm16", + "sample_rate": 16000, + "wait_null_answers": true, + "do_dialogue": false, + "do_punctuation": false, + "channelName": "channel_1" + } +} ``` -Если в config передать "do_dialogue":true и "do_punctuation":true то в последнем ответе будет предоставлен Капитализированный -текст, с пунктуацией разбитый на фразы. + +### Пример передачи данных (legacy) ```json - { - 'channel_name': 'Null', - 'silence': False, - 'data': { - 'result': - [ - {"conf": 1, "start": 0.04, "end": 0.36, "word": "ничьих"}, - {"conf": 1, "start": 0.52, "end": 0.56, "word": "не"}, - {"conf": 1, "start": 0.64, "end": 0.92, "word": "требуя" }, - {"conf": 1, "start": 1.08,"end": 1.44,"word": "похвал"}, - ], - "text": "ничьих не требуя похвал ... " - }, - 'error': None, - 'last_message': True, - 'sentenced_data': { - 'raw_text_sentenced_recognition': "channel_1: Ничьих, не требуя ... мои.\n channel_1: У Лукоморья дуб зеленый.", # текст построчно разбитый на фразы. - 'list_of_sentenced_recognitions': [{'start': 1.0, 'end': 1.28, 'text': 'У Лукоморья дуб зеленый.', 'speaker': 'channel_1'},... ] - "full_text_only": [ - "Ничьих, не требуя похвал. Счастлив уж я надеждой сладкой, что дева с трепетом любви посмотрит, может быть, украдкой на песни грешные мои. У Лукоморья дуб зеленый." - ], - } - } +{ + "text": "eof" +} +``` +### Ответы (legacy) + +```json +{ + "channel_name": "Null", + "silence": false, + "data": { + "result": [ + {"conf": 1.0, "start": 116.48, "end": 116.76, "word": "владимир"}, + {"conf": 1.0, "start": 116.92, "end": 117.48, "word": "анатольевич"} + ], + "text": "владимир анатольевич" + }, + "error": null, + "last_message": false, + "sentenced_data": {} +} ``` """ diff --git a/core/state_store.py b/core/state_store.py new file mode 100644 index 0000000..5b62049 --- /dev/null +++ b/core/state_store.py @@ -0,0 +1,97 @@ +""" +Модуль core/state_store.py +Абстракция хранилища состояния для WebSocket-сессий и мета-информации. +Подготовка к кластеризации (Redis) без изменения бизнес-логики. +""" + +import asyncio +from abc import ABC, abstractmethod +from typing import Any, Optional + + +class StateStore(ABC): + """ + Протокол (ABC) для хранилища состояния. + + Реализации должны поддерживать асинхронные операции get/set/delete + для сериализуемых мета-данных сессий (не AudioSegment). + """ + + @abstractmethod + async def get(self, key: str) -> Any: + """Возвращает значение по ключу или None.""" + + @abstractmethod + async def set(self, key: str, value: Any, ttl: Optional[int] = None) -> None: + """Сохраняет значение по ключу с опциональным TTL (секунды).""" + + @abstractmethod + async def delete(self, key: str) -> None: + """Удаляет ключ.""" + + @abstractmethod + async def exists(self, key: str) -> bool: + """Проверяет существование ключа.""" + + @abstractmethod + async def hgetall(self, key: str) -> dict: + """Возвращает dict, если значение — dict, иначе пустой dict.""" + + +class InMemoryStateStore(StateStore): + """ + Реализация StateStore в памяти текущего процесса (dict + asyncio.Lock). + + Хранит только сериализуемые мета-данные; AudioSegment-объекты + должны оставаться в AudioSession в памяти. + """ + + def __init__(self) -> None: + self._data: dict[str, Any] = {} + self._lock = asyncio.Lock() + + async def get(self, key: str) -> Any: + async with self._lock: + return self._data.get(key) + + async def set(self, key: str, value: Any, ttl: Optional[int] = None) -> None: + # TTL игнорируется в InMemory-реализации (нет фоновой очистки) + async with self._lock: + self._data[key] = value + + async def delete(self, key: str) -> None: + async with self._lock: + self._data.pop(key, None) + + async def exists(self, key: str) -> bool: + async with self._lock: + return key in self._data + + async def hgetall(self, key: str) -> dict: + val = await self.get(key) + if isinstance(val, dict): + return val + return {} + + +class RedisStateStore(StateStore): + """ + Заглушка для Redis-реализации StateStore. + + Полный набор сигнатур готов; реализация — в будущих этапах. + """ + + async def get(self, key: str) -> Any: + raise NotImplementedError("RedisStateStore.get is not implemented") + + async def set(self, key: str, value: Any, ttl: Optional[int] = None) -> None: + raise NotImplementedError("RedisStateStore.set is not implemented") + + async def delete(self, key: str) -> None: + raise NotImplementedError("RedisStateStore.delete is not implemented") + + async def exists(self, key: str) -> bool: + raise NotImplementedError("RedisStateStore.exists is not implemented") + + async def hgetall(self, key: str) -> dict: + raise NotImplementedError("RedisStateStore.hgetall is not implemented") diff --git a/main.py b/main.py index 8b7a126..048c461 100644 --- a/main.py +++ b/main.py @@ -90,6 +90,15 @@ async def lifespan(app): logger.debug("Приложение FastAPI запущено") app.state.start_time = time.time() + # Инициализация WebSocket-сервисов + from services.ws_manager import ConnectionManager + from services.ws_metrics import SystemMetricsCollector + from core.state_store import InMemoryStateStore + app.state.ws_manager = ConnectionManager(max_connections=settings.WS_MAX_CONNECTIONS) + app.state.metrics_collector = SystemMetricsCollector(start_time=app.state.start_time) + app.state.state_store = InMemoryStateStore() + logger.debug("WebSocket services initialized") + # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) @@ -129,6 +138,10 @@ async def lifespan(app): yield # Здесь приложение работает + # Graceful shutdown WebSocket (Задача 6.4, 6.9) + if hasattr(app.state, "ws_manager"): + await app.state.ws_manager.disconnect_all() + # cleanup (если нужно) if hasattr(app.state, "recognizer"): del app.state.recognizer diff --git a/models/ws_models.py b/models/ws_models.py index 8c4044c..a3cdb6e 100644 --- a/models/ws_models.py +++ b/models/ws_models.py @@ -32,6 +32,8 @@ class WSConfigMessage(WSBaseMessage): enable_diarization: bool = False num_speakers: int = Field(default=-1, ge=-1, le=10) enable_punctuation: bool = False + do_dialogue: bool = False + do_punctuation: bool = False audio_format: str = "pcm16" audio_transport: Literal["json_base64", "binary"] = Field( default="json_base64", diff --git a/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index 80d52c0..9d413f0 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -20,7 +20,8 @@ from Recognizer.engine.stream_recognition import simple_recognise from Punctuation import get_punctuator, SbertPuncCaseOnnx -#Todo Этот роут должен умереть. +# Todo: Этот роут — legacy. Удалить после полного перехода клиентов на /api/v1/asr/ws. +# Сохраняем "frozen" для обратной совместимости; не использовать в новых интеграциях. router = APIRouter() logger = logging.getLogger(__name__) diff --git a/services/asr_pipeline.py b/services/asr_pipeline.py new file mode 100644 index 0000000..d2fc602 --- /dev/null +++ b/services/asr_pipeline.py @@ -0,0 +1,264 @@ +""" +Модуль services/asr_pipeline.py +Содержит бизнес-логику потокового распознавания речи (ASR) через WebSocket: +накопление аудио, VAD-разделение по паузам, распознавание, постпроцессинг, +накопление результатов и формирование финального диалога с пунктуацией. +""" + +import logging +from io import BytesIO + +from pydub import AudioSegment + +from config import settings +from models.ws_models import ( + WSResultMessage, + WSRecognitionData, + WSWordItem, + WSErrorMessage, +) +from services.ws_session import AudioSession, SessionState +from services.ws_manager import ConnectionManager +from utils.resamppling import async_resample_audiosegment +from utils.tokens_to_Result import process_single_token_vocab_output +from utils.chunk_doing import find_last_speech_position_v2 +from Recognizer.engine.stream_recognition import simple_recognise +from Recognizer.engine.sentensizer import do_sensitizing + +logger = logging.getLogger(__name__) + + +async def process_audio_stream_chunk( + session: AudioSession, + chunk_bytes: bytes, + recognizer, + punctuator, + manager: ConnectionManager, +) -> None: + """ + Обрабатывает входящий чанк аудио в потоковом режиме. + + Алгоритм: + 1. Проверяет чётность байтов (дополняет до чётного при необходимости). + 2. Создаёт AudioSegment из raw bytes (PCM16) или через from_file для других форматов. + 3. Ресемплит до BASE_SAMPLE_RATE, приводит к моно. + 4. Накапливает в session.audio_buffer. + 5. При достижении MAX_OVERLAP_DURATION вызывает find_last_speech_position + (через временную синхронизацию с глобальными dict для сохранения точной логики VAD). + 6. Распознаёт готовый сегмент через simple_recognise. + 7. Применяет process_single_token_vocab_output со сдвигом времени. + 8. Отправляет результат клиенту (или silence partial при пустом тексте). + + Args: + session: Текущая аудио-сессия (содержит audio_buffer, audio_overlap и т.д.). + chunk_bytes: Сырые байты аудио от клиента. + recognizer: Экземпляр Recognizer. + punctuator: Экземпляр SbertPuncCaseOnnx (не используется в чанке, передаётся для единообразия). + manager: Менеджер WebSocket-соединений для отправки ответов. + """ + try: + # --- 1. Проверка чётности --- + if len(chunk_bytes) % 2 != 0: + chunk_bytes += bytes(2 - (len(chunk_bytes) % 2)) + + # --- 2. Создание AudioSegment --- + audio_format = session.config.audio_format if session.config else "pcm16" + sample_rate = session.config.sample_rate if session.config else settings.BASE_SAMPLE_RATE + + if audio_format == "pcm16": + audiosegment_chunk = AudioSegment( + chunk_bytes, + frame_rate=sample_rate, + sample_width=2, + channels=1, + ) + else: + buffer = BytesIO(chunk_bytes) + buffer.seek(0) + audiosegment_chunk = AudioSegment.from_file(buffer) + + # --- 3. Ресемплинг и моно --- + if audiosegment_chunk.frame_rate != settings.BASE_SAMPLE_RATE: + audiosegment_chunk = await async_resample_audiosegment(audiosegment_chunk, settings.BASE_SAMPLE_RATE) + if audiosegment_chunk.channels != 1: + audiosegment_chunk = audiosegment_chunk.set_channels(1) + + # --- 4. Накопление в буфер --- + session.audio_buffer += audiosegment_chunk + session.last_activity = __import__("time").time() + + # --- 5. Проверка порога VAD --- + combined_duration = (session.audio_overlap + session.audio_buffer).duration_seconds + if combined_duration < settings.MAX_OVERLAP_DURATION: + return + + # --- 5a. VAD-разделение через find_last_speech_position_v2 --- + try: + await find_last_speech_position_v2(session, is_last_chunk=False) + except Exception as exc: + logger.exception("VAD error for %s: %s", session.client_id, exc) + session.state = SessionState.error + error_msg = WSErrorMessage( + code="vad_error", + message=f"VAD processing error: {exc}", + is_fatal=False, + ) + try: + await manager.send_message(session.client_id, error_msg) + except Exception: + pass + return + + # --- 6. Распознавание последнего сегмента --- + if not session.audio_to_asr: + return + + segment = session.audio_to_asr[-1] + if segment.duration_seconds <= 0: + return + + asr_result = await simple_recognise(segment, recognizer=recognizer) + asr_result_words = process_single_token_vocab_output(asr_result, session.audio_duration) + session.audio_duration += segment.duration_seconds + + # Накопление для финального диалога + session.ws_collected_asr_res[f"channel_{1}"].append(asr_result_words) + + # --- 7. Отправка результата --- + text = asr_result_words.get("data", {}).get("text", "") + is_silence = len(text) == 0 or text == " " + + if is_silence: + if session.wait_null_answers: + msg = WSResultMessage( + channel_name=session.channel_name, + silence=True, + data=WSRecognitionData(), + error=None, + last_message=False, + ) + await manager.send_message(session.client_id, msg) + else: + logger.debug("Silence partial skipped (wait_null_answers=False)") + else: + words = asr_result_words.get("data", {}).get("result", []) + ws_words = [ + WSWordItem(conf=w["conf"], start=w["start"], end=w["end"], word=w["word"]) + for w in words + ] + data = WSRecognitionData(result=ws_words, text=text) + msg = WSResultMessage( + channel_name=session.channel_name, + silence=False, + data=data, + error=None, + last_message=False, + ) + await manager.send_message(session.client_id, msg) + + except Exception as exc: + logger.exception("ASR pipeline chunk error for %s: %s", session.client_id, exc) + session.state = SessionState.error + error_msg = WSErrorMessage( + code="asr_pipeline_error", + message=f"Chunk processing error: {exc}", + is_fatal=False, + ) + try: + await manager.send_message(session.client_id, error_msg) + except Exception: + pass + + +async def process_final_audio( + session: AudioSession, + recognizer, + punctuator, + manager: ConnectionManager, +) -> None: + """ + Обрабатывает финальный буфер аудио по получении EOF/EOS. + + Алгоритм: + 1. Объединяет audio_overlap + audio_buffer. + 2. Если длительность < 2 сек — дополняет тишиной до минимума. + 3. Распознаёт, постпроцессинг, добавляет в ws_collected_asr_res. + 4. Если do_dialogue — вызывает do_sensitizing с пунктуацией. + 5. Формирует и отправляет final WSResultMessage. + + Args: + session: Текущая аудио-сессия. + recognizer: Экземпляр Recognizer. + punctuator: Экземпляр SbertPuncCaseOnnx. + manager: Менеджер WebSocket-соединений. + """ + try: + # --- 1. Объединение остатков --- + final_audio = session.audio_overlap + session.audio_buffer + session.audio_to_asr.append(final_audio) + logger.debug("Final audio duration: %.3f sec", final_audio.duration_seconds) + + # --- 2. Дополнение тишиной при необходимости --- + if final_audio.duration_seconds < 2: + final_audio = final_audio + AudioSegment.silent(1000, frame_rate=settings.BASE_SAMPLE_RATE) + session.audio_to_asr[-1] = final_audio + logger.debug("Final audio padded with silence to %.3f sec", final_audio.duration_seconds) + + # --- 3. Распознавание --- + last_asr_result = await simple_recognise(final_audio, recognizer=recognizer) + last_result = process_single_token_vocab_output(last_asr_result, session.audio_duration) + session.ws_collected_asr_res[f"channel_{1}"].append(last_result) + + # --- 4. Определение silence --- + text = last_result.get("data", {}).get("text", "") + is_silence = len(text) == 0 or text == " " + if is_silence: + last_result = None + + # --- 5. Построение диалога / пунктуация --- + sentenced_data = None + if session.do_dialogue: + try: + sentenced_data = await do_sensitizing( + session.ws_collected_asr_res, + session.do_punctuation, + punctuator=punctuator, + ) + except Exception as exc: + logger.error("do_sensitizing error: %s", exc) + sentenced_data = None + + # --- 6. Формирование ответа --- + if is_silence: + data = WSRecognitionData() + else: + words = last_result.get("data", {}).get("result", []) + ws_words = [ + WSWordItem(conf=w["conf"], start=w["start"], end=w["end"], word=w["word"]) + for w in words + ] + data = WSRecognitionData(result=ws_words, text=text) + + final_msg = WSResultMessage( + channel_name=session.channel_name, + silence=is_silence, + data=data, + error=None, + last_message=True, + sentenced_data=sentenced_data, + ) + await manager.send_message(session.client_id, final_msg) + session.state = SessionState.completed + + except Exception as exc: + logger.exception("ASR pipeline final error for %s: %s", session.client_id, exc) + session.state = SessionState.error + error_msg = WSErrorMessage( + code="asr_pipeline_final_error", + message=f"Final processing error: {exc}", + is_fatal=False, + ) + try: + await manager.send_message(session.client_id, error_msg) + except Exception: + pass diff --git a/services/recognition_session.py b/services/recognition_session.py new file mode 100644 index 0000000..29f5545 --- /dev/null +++ b/services/recognition_session.py @@ -0,0 +1,161 @@ +""" +Модуль services/recognition_session.py +Базовый класс RecognitionSession и специализации для WebSocket и HTTP-файлового распознавания. +""" + +import os +import uuid +from enum import Enum +from typing import Optional, Any +from io import BytesIO + +from pydub import AudioSegment +from fastapi import UploadFile + +from config import settings +from models.ws_models import WSConfigMessage + + +class SessionState(str, Enum): + """Состояния жизненного цикла сессии распознавания.""" + created = "created" + processing = "processing" + completed = "completed" + error = "error" + + +class RecognitionSession: + """ + Базовая сессия распознавания речи. + Содержит общие поля для потокового (WS) и файлового (HTTP) распознавания: + AudioSegment-буферы, накопление результатов, конфигурация. + """ + + def __init__(self, session_id: Optional[str] = None) -> None: + """ + Инициализирует базовую сессию. + + Args: + session_id: UUID сессии. Если None — генерируется автоматически. + """ + self.session_id: str = session_id or str(uuid.uuid4()) + self.state: SessionState = SessionState.created + + # AudioSegment-буферы (общие для WS и файлового режима) + self.audio_buffer: AudioSegment = AudioSegment.silent( + 1, frame_rate=settings.BASE_SAMPLE_RATE + ) + self.audio_overlap: AudioSegment = AudioSegment.silent( + 1, frame_rate=settings.BASE_SAMPLE_RATE + ) + self.audio_to_asr: list[AudioSegment] = [] + self.audio_duration: float = 0.0 + + # Результаты и мета + self.collected_asr_res: dict = {f"channel_{1}": []} + self.channel_name: str = "Null" + self.do_dialogue: bool = False + self.do_punctuation: bool = False + self.config: Optional[WSConfigMessage] = None + + @property + def client_id(self) -> str: + """Возвращает session_id как client_id для совместимости с WS-обработчиками и VAD.""" + return self.session_id + + def reset(self) -> None: + """ + Очищает AudioSegment-буферы, результаты и сбрасывает состояние. + """ + self.audio_buffer = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + self.audio_overlap = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + self.audio_to_asr = [] + self.audio_duration = 0.0 + self.collected_asr_res = {f"channel_{1}": []} + self.state = SessionState.created + + def to_dict(self) -> dict: + """ + Сериализует мета-поля в dict для StateStore. + AudioSegment-объекты не сериализуются. + """ + return { + "session_id": self.session_id, + "state": self.state.value, + "audio_duration": self.audio_duration, + "channel_name": self.channel_name, + "do_dialogue": self.do_dialogue, + "do_punctuation": self.do_punctuation, + "collected_asr_res": self.collected_asr_res, + "config": self.config.model_dump() if self.config else None, + } + + @classmethod + def from_dict(cls, data: dict) -> "RecognitionSession": + """ + Восстанавливает сессию из dict (только мета-поля, без AudioSegment). + """ + session = cls(session_id=data.get("session_id")) + session.state = SessionState(data.get("state", "created")) + session.audio_duration = data.get("audio_duration", 0.0) + session.channel_name = data.get("channel_name", "Null") + session.do_dialogue = data.get("do_dialogue", False) + session.do_punctuation = data.get("do_punctuation", False) + session.collected_asr_res = data.get("collected_asr_res", {f"channel_{1}": []}) + if data.get("config"): + session.config = WSConfigMessage(**data["config"]) + return session + + +class FileRecognitionSession(RecognitionSession): + """ + Сессия для файлового распознавания (HTTP Upload / URL). + Расширяет RecognitionSession полями для работы с временными файлами. + """ + + def __init__( + self, + post_id: Optional[str] = None, + params: Optional[Any] = None, + session_id: Optional[str] = None, + ) -> None: + """ + Инициализирует файловую сессию. + + Args: + post_id: Идентификатор поста/задачи. + params: Параметры запроса (PostFileRequest или SyncASRRequest). + session_id: UUID сессии. + """ + super().__init__(session_id=session_id) + self.post_id: str = post_id or str(uuid.uuid4()) + self.params: Optional[Any] = params + self.tmp_path: Optional[str] = None + self.file_buffer: Optional[BytesIO] = None + + async def save_upload(self, file: UploadFile) -> None: + """ + Сохраняет загруженный файл во внутренний BytesIO. + + Args: + file: Объект UploadFile из FastAPI. + """ + self.file_buffer = BytesIO(await file.read()) + self.file_buffer.seek(0) + + def cleanup(self) -> None: + """ + Очищает ресурсы: закрывает BytesIO, удаляет временный файл с диска. + """ + if self.file_buffer is not None: + self.file_buffer.close() + self.file_buffer = None + if self.tmp_path and isinstance(self.tmp_path, (str, os.PathLike)): + try: + os.remove(self.tmp_path) + except OSError: + pass + self.tmp_path = None + elif self.tmp_path is not None: + # tmp_path содержит не строку/путь (например, BytesIO) — просто сбрасываем + self.tmp_path = None diff --git a/services/ws_session.py b/services/ws_session.py index 6a40a72..5252254 100644 --- a/services/ws_session.py +++ b/services/ws_session.py @@ -10,6 +10,7 @@ from typing import Optional import numpy as np +from pydub import AudioSegment from models.ws_models import WSConfigMessage from config import settings @@ -61,6 +62,21 @@ def __init__( ) self.last_activity: float = time.time() + # --- Поля для ASR pipeline (Задача 6.3) --- + self.audio_buffer: AudioSegment = AudioSegment.silent( + 1, frame_rate=settings.BASE_SAMPLE_RATE + ) + self.audio_overlap: AudioSegment = AudioSegment.silent( + 1, frame_rate=settings.BASE_SAMPLE_RATE + ) + self.audio_to_asr: list[AudioSegment] = [] + self.audio_duration: float = 0.0 + self.ws_collected_asr_res: dict = {f"channel_{1}": []} + self.wait_null_answers: bool = True + self.do_dialogue: bool = False + self.do_punctuation: bool = False + self.channel_name: str = "Null" + @property def current_buffer_duration_sec(self) -> float: """ @@ -132,6 +148,17 @@ async def reset(self) -> None: self.state = SessionState.connecting self.last_activity = time.time() + # Сброс ASR pipeline state + self.audio_buffer = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + self.audio_overlap = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + self.audio_to_asr = [] + self.audio_duration = 0.0 + self.ws_collected_asr_res = {f"channel_{1}": []} + self.wait_null_answers = True + self.do_dialogue = False + self.do_punctuation = False + self.channel_name = "Null" + def is_expired(self, timeout_sec: float) -> bool: """ Проверяет, истёк ли таймаут неактивности сессии. @@ -143,3 +170,46 @@ def is_expired(self, timeout_sec: float) -> bool: True, если с момента last_activity прошло больше timeout_sec. """ return (time.time() - self.last_activity) > timeout_sec + + def to_dict(self) -> dict: + """ + Сериализует лёгкие (не AudioSegment) поля сессии в dict для StateStore. + + Returns: + dict с client_id, config, channel_name, audio_duration, флагами + и накопленными результатами распознавания. + """ + return { + "client_id": self.client_id, + "config": self.config.model_dump() if self.config else None, + "channel_name": self.channel_name, + "audio_duration": self.audio_duration, + "do_dialogue": self.do_dialogue, + "do_punctuation": self.do_punctuation, + "wait_null_answers": self.wait_null_answers, + "ws_collected_asr_res": self.ws_collected_asr_res, + "state": self.state.value, + } + + @classmethod + def from_dict(cls, data: dict) -> "AudioSession": + """ + Восстанавливает сессию из dict (только мета-поля, без AudioSegment-буферов). + + Args: + data: dict, полученный из to_dict(). + + Returns: + AudioSession с восстановленной конфигурацией и флагами. + """ + session = cls(client_id=data["client_id"]) + if data.get("config"): + session.config = WSConfigMessage(**data["config"]) + session.channel_name = data.get("channel_name", "Null") + session.audio_duration = data.get("audio_duration", 0.0) + session.do_dialogue = data.get("do_dialogue", False) + session.do_punctuation = data.get("do_punctuation", False) + session.wait_null_answers = data.get("wait_null_answers", True) + session.ws_collected_asr_res = data.get("ws_collected_asr_res", {f"channel_{1}": []}) + session.state = SessionState(data.get("state", "connecting")) + return session diff --git a/templates/index.html b/templates/index.html index 5b6ab03..9e472aa 100644 --- a/templates/index.html +++ b/templates/index.html @@ -65,6 +65,10 @@

Распознать WAV файл. Обработка через WebSockets< Расстановка пунктуации +
+ + Оборачивать аудио в base64 (JSON). Если выключено — отправлять raw binary. +
@@ -252,6 +256,16 @@

Ответы сервера:

return { sampleRate, numChannels, bitsPerSample }; } + function arrayBufferToBase64(buffer) { + let binary = ''; + const bytes = new Uint8Array(buffer); + const len = bytes.byteLength; + for (let i = 0; i < len; i++) { + binary += String.fromCharCode(bytes[i]); + } + return window.btoa(binary); + } + async function sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate) { return new Promise((resolve, reject) => { const wsProtocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; @@ -262,18 +276,23 @@

Ответы сервера:

const doPunctuation = document.getElementById('ws_do_punctuation').checked; socket.onopen = async function() { + const useBase64 = document.getElementById('ws_use_base64').checked; + socket.send(JSON.stringify({ - config: { - sample_rate: sampleRate, - wait_null_answers: true, - do_dialogue: doDialogue, - do_punctuation: doPunctuation - } + type: "config", + sample_rate: sampleRate, + audio_format: "pcm16", + audio_transport: useBase64 ? "json_base64" : "binary", + wait_null_answers: true, + do_dialogue: doDialogue, + do_punctuation: doPunctuation, + channel_name: "channel_" + (channel + 1) })); const dataView = new Uint8Array(arrayBuffer); const bytesPerSample = 2; let offset = 44 + channel * bytesPerSample; + let seqNum = 0; while (offset < dataView.length) { const endOffset = Math.min(offset + chunkSize * numChannels, dataView.length); @@ -285,12 +304,21 @@

Ответы сервера:

chunk[i + 1] = dataView[sampleOffset + 1]; } - socket.send(chunk); + if (useBase64) { + const base64Chunk = arrayBufferToBase64(chunk); + socket.send(JSON.stringify({ + type: "audio_chunk", + audio_base64: base64Chunk, + seq_num: seqNum++ + })); + } else { + socket.send(chunk); + } offset += chunkSize * numChannels; await new Promise((resolve) => setTimeout(resolve, 100)); } - socket.send(JSON.stringify({ eof: 1 })); + socket.send(JSON.stringify({ type: "eos" })); }; socket.onmessage = function(event) { diff --git a/tests/test_asr_pipeline.py b/tests/test_asr_pipeline.py new file mode 100644 index 0000000..bff726f --- /dev/null +++ b/tests/test_asr_pipeline.py @@ -0,0 +1,233 @@ +""" +Тесты для services/asr_pipeline.py (потоковое распознавание с overlap). +""" + +import pytest +from unittest.mock import AsyncMock, patch, MagicMock +from pydub import AudioSegment + +from services.asr_pipeline import process_audio_stream_chunk, process_final_audio +from services.ws_session import AudioSession, SessionState +from services.ws_manager import ConnectionManager +from models.ws_models import WSConfigMessage, WSResultMessage, WSRecognitionData +from config import settings + + +class FakeManager(ConnectionManager): + """Фейковый менеджер для тестов ASR pipeline.""" + + def __init__(self): + super().__init__(max_connections=10) + self.sent_messages: list[tuple[str, object]] = [] + + async def send_message(self, client_id: str, message: object) -> None: + self.sent_messages.append((client_id, message)) + + +def _make_pcm16_silence(duration_sec: float, sample_rate: int = 16000) -> bytes: + """Генерирует PCM16 тишину заданной длительности.""" + samples = int(duration_sec * sample_rate) + return b"\x00\x00" * samples + + +@pytest.fixture +def session() -> AudioSession: + """Фикстура: AudioSession с типичной конфигурацией ASR.""" + sess = AudioSession(client_id="test-asr-42", max_buffer_duration_sec=300.0) + sess.config = WSConfigMessage( + sample_rate=16000, + audio_format="pcm16", + wait_null_answers=True, + do_dialogue=False, + do_punctuation=False, + channel_name="ch-test", + ) + sess.state = SessionState.receiving + return sess + + +@pytest.fixture +def manager() -> FakeManager: + return FakeManager() + + +@pytest.fixture +def recognizer(): + return MagicMock() + + +@pytest.fixture +def punctuator(): + return MagicMock() + + +class TestProcessAudioStreamChunk: + @pytest.mark.asyncio + async def test_small_chunk_accumulates_no_vad(self, session, manager, recognizer, punctuator): + """Малый чанк (< MAX_OVERLAP_DURATION) накапливается, VAD не вызывается.""" + chunk = _make_pcm16_silence(1.0) # 1 сек < MAX_OVERLAP_DURATION (обычно 30) + with patch("services.asr_pipeline.find_last_speech_position_v2", new_callable=AsyncMock) as mock_vad: + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + await process_audio_stream_chunk(session, chunk, recognizer, punctuator, manager) + mock_vad.assert_not_awaited() + mock_asr.assert_not_awaited() + assert session.audio_buffer.duration_seconds >= 1.0 + + @pytest.mark.asyncio + async def test_large_chunk_triggers_vad_and_sends_result(self, session, manager, recognizer, punctuator): + """Большой чанк вызывает VAD, распознавание и отправку результата.""" + # Устанавливаем буфер так, чтобы combined_duration >= MAX_OVERLAP_DURATION + session.audio_buffer = AudioSegment.silent( + int((settings.MAX_OVERLAP_DURATION - 0.2) * 1000), frame_rate=16000 + ) + chunk = _make_pcm16_silence(0.5) + + async def mock_vad(session, is_last_chunk): + session.audio_to_asr.append(session.audio_buffer) + session.audio_overlap = AudioSegment.silent(1, frame_rate=16000) + session.audio_buffer = AudioSegment.silent(1, frame_rate=16000) + + with patch("services.asr_pipeline.find_last_speech_position_v2", mock_vad): + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + mock_asr.return_value = { + "tokens": ["п", "р", "и", "в", "е", "т"], + "timestamps": [0.0, 0.1, 0.2, 0.3, 0.4, 0.5], + "text": "привет", + } + await process_audio_stream_chunk(session, chunk, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + client_id, msg = manager.sent_messages[0] + assert client_id == session.client_id + assert isinstance(msg, WSResultMessage) + assert msg.silence is False + assert msg.data.text == "привет" + assert msg.last_message is False + + @pytest.mark.asyncio + async def test_silence_partial_when_empty_text(self, session, manager, recognizer, punctuator): + """При пустом тексте и wait_null_answers=True отправляется silence partial.""" + session.audio_buffer = AudioSegment.silent( + int((settings.MAX_OVERLAP_DURATION - 0.05) * 1000), frame_rate=16000 + ) + chunk = _make_pcm16_silence(0.1) + + async def mock_vad(session, is_last_chunk): + session.audio_to_asr.append(session.audio_buffer) + session.audio_overlap = AudioSegment.silent(1, frame_rate=16000) + session.audio_buffer = AudioSegment.silent(1, frame_rate=16000) + + with patch("services.asr_pipeline.find_last_speech_position_v2", mock_vad): + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + mock_asr.return_value = { + "tokens": [], + "timestamps": [], + "text": "", + } + await process_audio_stream_chunk(session, chunk, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert isinstance(msg, WSResultMessage) + assert msg.silence is True + assert msg.data.text == "" + + @pytest.mark.asyncio + async def test_odd_bytes_padded(self, session, manager, recognizer, punctuator): + """Нечётное количество байтов дополняется до чётного.""" + chunk = b"\x01\x02\x03" # 3 байта + with patch("services.asr_pipeline.find_last_speech_position_v2", new_callable=AsyncMock): + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock): + await process_audio_stream_chunk(session, chunk, recognizer, punctuator, manager) + assert session.audio_buffer.duration_seconds > 0 + + @pytest.mark.asyncio + async def test_error_sends_ws_error_message(self, session, manager, recognizer, punctuator): + """При исключении внутри pipeline отправляется WSErrorMessage.""" + chunk = _make_pcm16_silence(1.0) + session.audio_buffer = AudioSegment.silent( + int((settings.MAX_OVERLAP_DURATION - 0.2) * 1000), frame_rate=16000 + ) + with patch("services.asr_pipeline.find_last_speech_position_v2", side_effect=ValueError("boom")): + await process_audio_stream_chunk(session, chunk, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert msg.type == "error" + assert msg.code == "vad_error" + + +class TestProcessFinalAudio: + @pytest.mark.asyncio + async def test_final_under_2sec_padded(self, session, manager, recognizer, punctuator): + """Финальный аудио < 2 сек дополняется тишиной и распознаётся.""" + session.audio_overlap = AudioSegment.silent(500, frame_rate=16000) # 0.5 сек + session.audio_buffer = AudioSegment.silent(500, frame_rate=16000) # 0.5 сек + + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + mock_asr.return_value = { + "tokens": ["т", "е", "с", "т"], + "timestamps": [0.0, 0.1, 0.2, 0.3], + "text": "тест", + } + await process_final_audio(session, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert isinstance(msg, WSResultMessage) + assert msg.last_message is True + assert msg.silence is False + assert msg.data.text == "тест" + + @pytest.mark.asyncio + async def test_final_empty_silence(self, session, manager, recognizer, punctuator): + """Финальный результат пустой — отправляется silence=True.""" + session.audio_overlap = AudioSegment.silent(1000, frame_rate=16000) + session.audio_buffer = AudioSegment.silent(1000, frame_rate=16000) + + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + mock_asr.return_value = { + "tokens": [], + "timestamps": [], + "text": "", + } + await process_final_audio(session, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert msg.silence is True + assert msg.last_message is True + + @pytest.mark.asyncio + async def test_final_with_dialogue(self, session, manager, recognizer, punctuator): + """При do_dialogue=True вызывается do_sensitizing и sentenced_data заполняется.""" + session.do_dialogue = True + session.do_punctuation = True + session.audio_overlap = AudioSegment.silent(2000, frame_rate=16000) + session.audio_buffer = AudioSegment.silent(1000, frame_rate=16000) + + with patch("services.asr_pipeline.simple_recognise", new_callable=AsyncMock) as mock_asr: + mock_asr.return_value = { + "tokens": ["п", "р", "и", "в", "е", "т"], + "timestamps": [0.0, 0.1, 0.2, 0.3, 0.4, 0.5], + "text": "привет", + } + with patch("services.asr_pipeline.do_sensitizing", new_callable=AsyncMock) as mock_sent: + mock_sent.return_value = {"raw_text_sentenced_recognition": "Ch1: Привет."} + await process_final_audio(session, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert msg.sentenced_data == {"raw_text_sentenced_recognition": "Ch1: Привет."} + + @pytest.mark.asyncio + async def test_final_error_sends_ws_error(self, session, manager, recognizer, punctuator): + """При исключении на финальном этапе отправляется WSErrorMessage.""" + with patch("services.asr_pipeline.simple_recognise", side_effect=RuntimeError("boom")): + await process_final_audio(session, recognizer, punctuator, manager) + + assert len(manager.sent_messages) == 1 + _, msg = manager.sent_messages[0] + assert msg.type == "error" + assert msg.code == "asr_pipeline_final_error" + assert session.state == SessionState.error diff --git a/tests/test_chunk_doing_v2.py b/tests/test_chunk_doing_v2.py new file mode 100644 index 0000000..6e63e70 --- /dev/null +++ b/tests/test_chunk_doing_v2.py @@ -0,0 +1,47 @@ +""" +Тесты для find_last_speech_position_v2 (utils/chunk_doing.py). +Проверяет корректность VAD-разделения с фиктивными AudioSegment. +""" + +import pytest +from unittest.mock import AsyncMock, patch, MagicMock +from pydub import AudioSegment + +from utils.chunk_doing import find_last_speech_position_v2 +from services.ws_session import AudioSession + + +@pytest.fixture +def session() -> AudioSession: + sess = AudioSession(client_id="vad-test", max_buffer_duration_sec=300.0) + sess.audio_buffer = AudioSegment.silent(5000, frame_rate=16000) # 5 сек + sess.audio_overlap = AudioSegment.silent(1, frame_rate=16000) + sess.audio_to_asr = [] + return sess + + +class TestFindLastSpeechPositionV2: + @pytest.mark.asyncio + async def test_last_chunk_splits_to_asr(self, session: AudioSession) -> None: + """При is_last_chunk=True весь буфер уходит в audio_to_asr по частям.""" + with patch("utils.chunk_doing.vad", MagicMock()): + await find_last_speech_position_v2(session, is_last_chunk=True) + assert len(session.audio_to_asr) > 0 + # Проверяем, что суммарная длительность равна исходной + total = sum(seg.duration_seconds for seg in session.audio_to_asr) + assert total == pytest.approx(5.0, 0.1) + + @pytest.mark.asyncio + async def test_non_last_chunk_clears_buffer(self, session: AudioSession) -> None: + """При is_last_chunk=False буфер очищается, overlap получает хвост.""" + with patch("utils.chunk_doing.vad", MagicMock()) as mock_vad: + mock_vad.reset_state = AsyncMock() + mock_vad.state = None + mock_vad.is_speech = AsyncMock(return_value=(0.99, None)) # всегда речь + mock_vad.prob_level = 0.5 + + await find_last_speech_position_v2(session, is_last_chunk=False) + # Если весь сегмент — речь, всё уходит в audio_to_asr + assert len(session.audio_to_asr) == 1 + assert session.audio_buffer.duration_seconds < 0.1 # silent(1ms) + assert session.audio_overlap.duration_seconds == pytest.approx(0.0, 0.1) diff --git a/tests/test_recognition_session.py b/tests/test_recognition_session.py new file mode 100644 index 0000000..8577f6d --- /dev/null +++ b/tests/test_recognition_session.py @@ -0,0 +1,96 @@ +""" +Тесты для services/recognition_session.py. +""" + +import io +import os + +import pytest +from unittest.mock import MagicMock, AsyncMock + +from services.recognition_session import RecognitionSession, FileRecognitionSession, SessionState +from models.ws_models import WSConfigMessage + + +class TestRecognitionSession: + def test_init_generates_uuid(self): + """Проверяет автоматическую генерацию session_id.""" + sess = RecognitionSession() + assert len(sess.session_id) == 36 + assert sess.state == SessionState.created + + def test_audio_buffers_are_silent_segment(self): + """Проверяет, что буферы инициализируются как AudioSegment.silent.""" + sess = RecognitionSession() + assert sess.audio_buffer.duration_seconds < 0.1 + assert sess.audio_overlap.duration_seconds < 0.1 + assert sess.audio_to_asr == [] + + def test_reset_clears_buffers(self): + """Проверяет очистку буферов и сброс состояния.""" + sess = RecognitionSession() + sess.audio_to_asr = [MagicMock()] + sess.audio_duration = 10.0 + sess.reset() + assert sess.audio_to_asr == [] + assert sess.audio_duration == 0.0 + assert sess.state == SessionState.created + + def test_to_dict_excludes_audio_segment(self): + """Проверяет, что to_dict сериализует только мета-поля.""" + sess = RecognitionSession() + sess.config = WSConfigMessage(sample_rate=16000, do_dialogue=True) + sess.do_dialogue = True + sess.audio_duration = 5.0 + d = sess.to_dict() + assert d["session_id"] == sess.session_id + assert d["audio_duration"] == 5.0 + assert d["do_dialogue"] is True + assert "audio_buffer" not in d + + def test_from_dict_restores_meta(self): + """Проверяет восстановление сессии из dict.""" + sess = RecognitionSession() + sess.config = WSConfigMessage(sample_rate=8000, do_punctuation=True) + sess.channel_name = "ch-1" + d = sess.to_dict() + restored = RecognitionSession.from_dict(d) + assert restored.channel_name == "ch-1" + assert restored.config.sample_rate == 8000 + assert restored.state == SessionState.created + + +class TestFileRecognitionSession: + def test_init_with_post_id(self): + """Проверяет инициализацию с кастомным post_id и params.""" + sess = FileRecognitionSession(post_id="post-123", params={"foo": "bar"}) + assert sess.post_id == "post-123" + assert sess.params == {"foo": "bar"} + + @pytest.mark.asyncio + async def test_save_upload(self): + """Проверяет сохранение UploadFile в BytesIO.""" + sess = FileRecognitionSession() + mock_file = MagicMock() + mock_file.read = AsyncMock(return_value=b"audio_data") + await sess.save_upload(mock_file) + assert sess.file_buffer is not None + assert sess.file_buffer.read() == b"audio_data" + + def test_cleanup_closes_buffer(self): + """Проверяет закрытие BytesIO при cleanup.""" + sess = FileRecognitionSession() + sess.file_buffer = io.BytesIO(b"data") + sess.tmp_path = None + sess.cleanup() + assert sess.file_buffer is None + + def test_cleanup_removes_tmp_path(self, tmp_path): + """Проверяет удаление временного файла с диска.""" + fake_file = tmp_path / "test.wav" + fake_file.write_text("fake audio") + sess = FileRecognitionSession() + sess.tmp_path = str(fake_file) + sess.cleanup() + assert not os.path.exists(str(fake_file)) + assert sess.tmp_path is None diff --git a/tests/test_state_store.py b/tests/test_state_store.py new file mode 100644 index 0000000..de53aac --- /dev/null +++ b/tests/test_state_store.py @@ -0,0 +1,92 @@ +""" +Тесты для core/state_store.py (InMemoryStateStore и RedisStateStore-заглушка). +""" + +import asyncio + +import pytest + +from core.state_store import InMemoryStateStore, RedisStateStore + + +class TestInMemoryStateStore: + @pytest.mark.asyncio + async def test_set_and_get(self): + store = InMemoryStateStore() + await store.set("key1", "value1") + assert await store.get("key1") == "value1" + + @pytest.mark.asyncio + async def test_get_missing_returns_none(self): + store = InMemoryStateStore() + assert await store.get("missing") is None + + @pytest.mark.asyncio + async def test_delete(self): + store = InMemoryStateStore() + await store.set("key1", "value1") + await store.delete("key1") + assert await store.get("key1") is None + + @pytest.mark.asyncio + async def test_exists(self): + store = InMemoryStateStore() + await store.set("key1", "value1") + assert await store.exists("key1") is True + assert await store.exists("key2") is False + + @pytest.mark.asyncio + async def test_hgetall_dict(self): + store = InMemoryStateStore() + await store.set("hash1", {"a": 1, "b": 2}) + assert await store.hgetall("hash1") == {"a": 1, "b": 2} + + @pytest.mark.asyncio + async def test_hgetall_non_dict(self): + store = InMemoryStateStore() + await store.set("key1", "string") + assert await store.hgetall("key1") == {} + + @pytest.mark.asyncio + async def test_concurrent_access(self): + store = InMemoryStateStore() + await store.set("counter", 0) + # Инкремент в 10 корутинах + async def inc(): + val = await store.get("counter") + await store.set("counter", val + 1) + + await asyncio.gather(*[inc() for _ in range(10)]) + assert await store.get("counter") == 10 + + +class TestRedisStateStore: + @pytest.mark.asyncio + async def test_get_raises_not_implemented(self): + store = RedisStateStore() + with pytest.raises(NotImplementedError): + await store.get("key") + + @pytest.mark.asyncio + async def test_set_raises_not_implemented(self): + store = RedisStateStore() + with pytest.raises(NotImplementedError): + await store.set("key", "val") + + @pytest.mark.asyncio + async def test_delete_raises_not_implemented(self): + store = RedisStateStore() + with pytest.raises(NotImplementedError): + await store.delete("key") + + @pytest.mark.asyncio + async def test_exists_raises_not_implemented(self): + store = RedisStateStore() + with pytest.raises(NotImplementedError): + await store.exists("key") + + @pytest.mark.asyncio + async def test_hgetall_raises_not_implemented(self): + store = RedisStateStore() + with pytest.raises(NotImplementedError): + await store.hgetall("key") diff --git a/tests/test_ws_models.py b/tests/test_ws_models.py index 0aa8a1b..4e316f4 100644 --- a/tests/test_ws_models.py +++ b/tests/test_ws_models.py @@ -35,6 +35,12 @@ def test_audio_transport_binary(self): msg = WSConfigMessage.model_validate_json(raw) assert msg.audio_transport == "binary" + def test_do_dialogue_and_punctuation(self): + raw = '{"type": "config", "do_dialogue": true, "do_punctuation": true}' + msg = WSConfigMessage.model_validate_json(raw) + assert msg.do_dialogue is True + assert msg.do_punctuation is True + def test_from_json(self): raw = '{"type": "config", "sample_rate": 8000, "enable_punctuation": true}' msg = WSConfigMessage.model_validate_json(raw) diff --git a/tests/test_ws_session.py b/tests/test_ws_session.py index bcd43a9..c76c3f6 100644 --- a/tests/test_ws_session.py +++ b/tests/test_ws_session.py @@ -112,3 +112,30 @@ async def test_duration_respects_config_sample_rate(self): sess.config = WSConfigMessage(sample_rate=8000) await sess.add_audio(np.zeros(8000, dtype=np.float32)) # 1 сек при 8кГц assert sess.current_buffer_duration_sec == pytest.approx(1.0, 0.01) + + +class TestAudioSessionSerialization: + def test_to_dict_and_from_dict(self): + """Сериализация и восстановление мета-полей сессии.""" + sess = AudioSession(client_id="ser-test") + sess.config = WSConfigMessage(sample_rate=8000, do_dialogue=True, do_punctuation=True) + sess.channel_name = "ch-1" + sess.audio_duration = 5.0 + sess.ws_collected_asr_res = {"channel_1": [{"text": "привет"}]} + sess.do_dialogue = True + sess.do_punctuation = True + + d = sess.to_dict() + assert d["client_id"] == "ser-test" + assert d["channel_name"] == "ch-1" + assert d["audio_duration"] == 5.0 + assert d["do_dialogue"] is True + assert d["do_punctuation"] is True + + restored = AudioSession.from_dict(d) + assert restored.client_id == "ser-test" + assert restored.channel_name == "ch-1" + assert restored.audio_duration == 5.0 + assert restored.do_dialogue is True + assert restored.config.sample_rate == 8000 + assert restored.ws_collected_asr_res == {"channel_1": [{"text": "привет"}]} From 03a922fc6a6592e8cc8073bc59255ea89784639f Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 7 May 2026 15:57:25 +0300 Subject: [PATCH 28/46] =?UTF-8?q?=D0=92=D1=8B=D0=B2=D0=BE=D0=B4=D0=B8?= =?UTF-8?q?=D0=BC=20gl=D0=BEbals=20=D0=B8=D0=B7=20url=20=D0=B8=20file.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/engine/file_recognition.py | 9 ++- api/legacy/is_alive.py | 31 +++++--- api/legacy/post_by_url.py | 7 +- api/v1/endpoints/asr_file.py | 62 ++++++--------- api/v1/endpoints/asr_url.py | 46 +++++++---- api/v1/endpoints/asr_ws.py | 3 +- services/recognition_session.py | 4 +- services/ws_session.py | 87 ++++++++------------ utils/chunk_doing.py | 110 +++++++++++++++++++++++++- utils/get_audio_file.py | 27 +++---- 10 files changed, 242 insertions(+), 144 deletions(-) diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index a31ea33..ad4ec17 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -39,10 +39,13 @@ def process_file(session=None, recognizer=None, punctuator=None, diarizer=None, logger.debug(f'Принят новый "post_file" id = {post_id}') try: + file_source = session.tmp_path if session.tmp_path is not None else session.file_buffer + if file_source is None: + raise ValueError("FileRecognitionSession has no tmp_path or file_buffer") if params.make_mono: - audio_input = AudioSegment.from_file(session.tmp_path).set_channels(1) + audio_input = AudioSegment.from_file(file_source).set_channels(1) else: - audio_input = AudioSegment.from_file(session.tmp_path) + audio_input = AudioSegment.from_file(file_source) except Exception as e: error_description += f"Error loading audio file: {e}" logger.error(error_description) @@ -159,7 +162,7 @@ def process_file(session=None, recognizer=None, punctuator=None, diarizer=None, # Проверяем возможность диаризации. Если здесь стерео-канал, то диаризацию выключаем. elif params.do_diarization and len(audio_input.split_to_mono()) != 1: error_description += f"Only mono diarization available.\n" - logger.warn("При запрошенной диаризации аудио имеет более одного аудио-канала. Диаризация будет выключена.") + logger.warning("При запрошенной диаризации аудио имеет более одного аудио-канала. Диаризация будет выключена.") params.do_diarization = False if params.do_diarization: diff --git a/api/legacy/is_alive.py b/api/legacy/is_alive.py index ffcdaee..a2480dc 100644 --- a/api/legacy/is_alive.py +++ b/api/legacy/is_alive.py @@ -1,16 +1,28 @@ from fastapi import APIRouter import logging import pynvml -from utils.pre_start_init import audio_to_asr logger = logging.getLogger(__name__) router = APIRouter() +# Кэшированный handle pynvml для избежания Init/Shutdown на каждый запрос +_nvml_handle = None + +def _get_nvml_handle(): + global _nvml_handle + if _nvml_handle is None: + try: + pynvml.nvmlInit() + _nvml_handle = pynvml.nvmlDeviceGetHandleByIndex(0) + except Exception: + return None + return _nvml_handle + def get_gpu_free_memory(): + handle = _get_nvml_handle() + if handle is None: + return {"error": "GPU not available or pynvml not initialized"}, None, None, None try: - # todo Доработать, выдавать ответ в зависимости от провайдера. - pynvml.nvmlInit() - handle = pynvml.nvmlDeviceGetHandleByIndex(0) # Первая видеокарта mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle) free_mb = mem_info.free / 1024**2 utilization = pynvml.nvmlDeviceGetUtilizationRates(handle) @@ -18,9 +30,7 @@ def get_gpu_free_memory(): temperature = pynvml.nvmlDeviceGetTemperature(handle, pynvml.NVML_TEMPERATURE_GPU) except pynvml.NVMLError as e: return {"error": str(e)}, None, None, None - finally: - pynvml.nvmlShutdown() - return None, free_mb, gpu_load,temperature + return None, free_mb, gpu_load, temperature @router.get("/is_alive") @@ -30,11 +40,12 @@ async def check_if_service_is_alive(): ) error_description = None logging.info('GET_is_alive') - tasks_in_work = len(audio_to_asr) + # Legacy: глобальный audio_to_asr больше не используется в новой архитектуре + tasks_in_work = 0 - error, free_mb, gpu_load,temperature = get_gpu_free_memory() + error, free_mb, gpu_load, temperature = get_gpu_free_memory() if error: - error_description = error.get("error",None) + error_description = error.get("error", None) if tasks_in_work == 0: state = "idle" diff --git a/api/legacy/post_by_url.py b/api/legacy/post_by_url.py index 5da6069..1fe4440 100644 --- a/api/legacy/post_by_url.py +++ b/api/legacy/post_by_url.py @@ -1,7 +1,6 @@ import uuid import asyncio from fastapi import APIRouter, Depends -from utils.pre_start_init import posted_and_downloaded_audio from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, BaseResponse @@ -46,9 +45,9 @@ async def post(params: SyncASRRequest, # Получаем файл post_id = uuid.uuid4() if params.AudioFileUrl: - res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) + res, error_description, buffer = await getting_audiofile(params.AudioFileUrl, post_id) else: - res, error_description = await open_default_audiofile(post_id) + res, error_description, buffer = await open_default_audiofile(post_id) if not res: logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') @@ -63,7 +62,7 @@ async def post(params: SyncASRRequest, try: # Запускаем обработку в потоке result_dict = await asyncio.to_thread(process_file, - tmp_path=posted_and_downloaded_audio[post_id], + tmp_path=buffer, params=params, recognizer=recognizer, punctuator=punctuator, diff --git a/api/v1/endpoints/asr_file.py b/api/v1/endpoints/asr_file.py index 990fd58..43dea75 100644 --- a/api/v1/endpoints/asr_file.py +++ b/api/v1/endpoints/asr_file.py @@ -1,9 +1,9 @@ -from io import BytesIO import asyncio import logging from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile from models.fast_api_models import PostFileRequest, V1BaseResponse, ASRData, RawData, SentencedData, DiarizedData +from services.recognition_session import FileRecognitionSession from Recognizer.engine.file_recognition import process_file from Recognizer import get_recognizer, Recognizer from Punctuation import get_punctuator, SbertPuncCaseOnnx @@ -52,45 +52,35 @@ async def async_receive_file( punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), diarizer: Diarizer = Depends(get_diarizer) ) -> V1BaseResponse: + session = FileRecognitionSession(params=params) try: - buffer = BytesIO(await file.read()) - buffer.seek(0) + await session.save_upload(file) + logger.info(f"Получен и сохранён файл {file.filename}") + result_dict = await asyncio.to_thread( + process_file, + session=session, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer + ) + return V1BaseResponse( + success=result_dict.get('success', True), + error_description=result_dict.get('error_description'), + data=ASRData( + raw_data=RawData.model_validate(result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, + sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, + diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None + ) + ) except Exception as e: - error_description = f"Не удалось сохранить файл для распознавания: {file.filename}, размер файла: {file.size}, по причине: {e}" + error_description = f"Ошибка обработки в process_file - {e}" logger.error(error_description) return V1BaseResponse( success=False, - error_description=error_description, + error_description=str(error_description), data=ASRData() ) - else: - logger.info(f"Получен и сохранён файл {file.filename}") - try: - result_dict = await asyncio.to_thread( - process_file, - tmp_path=buffer, - params=params, - recognizer=recognizer, - punctuator=punctuator, - diarizer=diarizer - ) - return V1BaseResponse( - success=result_dict.get('success', True), - error_description=result_dict.get('error_description'), - data=ASRData( - raw_data=RawData.model_validate(result_dict.get('raw_data', {})) if result_dict.get('raw_data') else None, - sentenced_data=SentencedData(**result_dict.get('sentenced_data', {})) if result_dict.get('sentenced_data') else None, - diarized_data=DiarizedData(**result_dict.get('diarized_data', {})) if result_dict.get('diarized_data') else None - ) - ) - except Exception as e: - error_description = f"Ошибка обработки в process_file - {e}" - logger.error(error_description) - return V1BaseResponse( - success=False, - error_description=str(error_description), - data=ASRData() - ) - finally: - await file.close() - del file + finally: + await file.close() + session.cleanup() + await session.reset() diff --git a/api/v1/endpoints/asr_url.py b/api/v1/endpoints/asr_url.py index 490a476..0b83a6c 100644 --- a/api/v1/endpoints/asr_url.py +++ b/api/v1/endpoints/asr_url.py @@ -1,9 +1,9 @@ import uuid import asyncio import logging -from utils.pre_start_init import posted_and_downloaded_audio from utils.get_audio_file import getting_audiofile, open_default_audiofile from models.fast_api_models import SyncASRRequest, V1BaseResponse, ASRData, RawData, SentencedData, DiarizedData +from services.recognition_session import FileRecognitionSession from fastapi import APIRouter, Depends from Recognizer import get_recognizer, Recognizer @@ -25,26 +25,35 @@ async def post_v1( diarizer: Diarizer = Depends(get_diarizer) ) -> V1BaseResponse: post_id = uuid.uuid4() - if params.AudioFileUrl: - res, error_description = await getting_audiofile(params.AudioFileUrl, post_id) - else: - res, error_description = await open_default_audiofile(post_id) + session = FileRecognitionSession(post_id=str(post_id), params=params) + try: + if params.AudioFileUrl: + res, error_description, file_buffer = await getting_audiofile(params.AudioFileUrl, post_id) + else: + res, error_description, file_buffer = await open_default_audiofile(post_id) - if not res: - logger.error( - f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}' - ) - return V1BaseResponse( - success=False, - error_description=error_description, - data=ASRData() - ) + if not res: + logger.error( + f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}' + ) + return V1BaseResponse( + success=False, + error_description=error_description, + data=ASRData() + ) + + if file_buffer is None: + return V1BaseResponse( + success=False, + error_description="Получен пустой буфер файла", + data=ASRData() + ) + session.file_buffer = file_buffer + session.file_buffer.seek(0) - try: result_dict = await asyncio.to_thread( process_file, - tmp_path=posted_and_downloaded_audio[post_id], - params=params, + session=session, recognizer=recognizer, punctuator=punctuator, diarizer=diarizer @@ -66,3 +75,6 @@ async def post_v1( error_description=str(error_description), data=ASRData() ) + finally: + session.cleanup() + await session.reset() diff --git a/api/v1/endpoints/asr_ws.py b/api/v1/endpoints/asr_ws.py index 72abc9d..a0dec23 100644 --- a/api/v1/endpoints/asr_ws.py +++ b/api/v1/endpoints/asr_ws.py @@ -38,14 +38,13 @@ async def audio_session_lifecycle(client_id: str): """ Контекстный менеджер жизненного цикла AudioSession. - Гарантирует очистку AudioSegment-буферов и глобальных dict при выходе. + Гарантирует очистку AudioSegment-буферов. """ session = AudioSession(client_id=client_id) try: yield session finally: await session.reset() - _cleanup_globals(client_id) @router.websocket("/ws") diff --git a/services/recognition_session.py b/services/recognition_session.py index 29f5545..063fea7 100644 --- a/services/recognition_session.py +++ b/services/recognition_session.py @@ -19,6 +19,8 @@ class SessionState(str, Enum): """Состояния жизненного цикла сессии распознавания.""" created = "created" + connecting = "connecting" + receiving = "receiving" processing = "processing" completed = "completed" error = "error" @@ -63,7 +65,7 @@ def client_id(self) -> str: """Возвращает session_id как client_id для совместимости с WS-обработчиками и VAD.""" return self.session_id - def reset(self) -> None: + async def reset(self) -> None: """ Очищает AudioSegment-буферы, результаты и сбрасывает состояние. """ diff --git a/services/ws_session.py b/services/ws_session.py index 5252254..e2d8729 100644 --- a/services/ws_session.py +++ b/services/ws_session.py @@ -6,26 +6,16 @@ import time from collections import deque -from enum import Enum from typing import Optional import numpy as np -from pydub import AudioSegment from models.ws_models import WSConfigMessage +from services.recognition_session import RecognitionSession, SessionState from config import settings -class SessionState(str, Enum): - """Состояния жизненного цикла аудио-сессии.""" - connecting = "connecting" - receiving = "receiving" - processing = "processing" - completed = "completed" - error = "error" - - -class AudioSession: +class AudioSession(RecognitionSession): """ Управляет буфером аудио и состоянием для одного WebSocket-клиента. @@ -51,31 +41,25 @@ def __init__( max_buffer_duration_sec: Максимальная длительность буфера (сек). По умолчанию берётся из settings.WS_MAX_BUFFER_DURATION_SEC. """ - self.client_id: str = client_id - self.state: SessionState = SessionState.connecting - self.buffer: deque[np.ndarray] = deque() - self.config: Optional[WSConfigMessage] = None + super().__init__(session_id=client_id) + self.state = SessionState.connecting + self.wait_null_answers: bool = True + self.last_activity: float = time.time() self.max_buffer_duration_sec: float = ( max_buffer_duration_sec if max_buffer_duration_sec is not None else getattr(settings, "WS_MAX_BUFFER_DURATION_SEC", 300.0) ) - self.last_activity: float = time.time() + self.buffer: deque[np.ndarray] = deque() - # --- Поля для ASR pipeline (Задача 6.3) --- - self.audio_buffer: AudioSegment = AudioSegment.silent( - 1, frame_rate=settings.BASE_SAMPLE_RATE - ) - self.audio_overlap: AudioSegment = AudioSegment.silent( - 1, frame_rate=settings.BASE_SAMPLE_RATE - ) - self.audio_to_asr: list[AudioSegment] = [] - self.audio_duration: float = 0.0 - self.ws_collected_asr_res: dict = {f"channel_{1}": []} - self.wait_null_answers: bool = True - self.do_dialogue: bool = False - self.do_punctuation: bool = False - self.channel_name: str = "Null" + @property + def ws_collected_asr_res(self) -> dict: + """Alias для collected_asr_res из RecognitionSession (обратная совместимость).""" + return self.collected_asr_res + + @ws_collected_asr_res.setter + def ws_collected_asr_res(self, value: dict) -> None: + self.collected_asr_res = value @property def current_buffer_duration_sec(self) -> float: @@ -140,24 +124,13 @@ async def get_full_audio(self) -> np.ndarray: async def reset(self) -> None: """ - Очищает буфер, сбрасывает конфигурацию и переводит сессию - в начальное состояние (connecting). + Очищает WS-специфичные буферы и сбрасывает базовое состояние. """ - self.buffer.clear() - self.config = None + await super().reset() self.state = SessionState.connecting + self.buffer.clear() self.last_activity = time.time() - - # Сброс ASR pipeline state - self.audio_buffer = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - self.audio_overlap = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) - self.audio_to_asr = [] - self.audio_duration = 0.0 - self.ws_collected_asr_res = {f"channel_{1}": []} self.wait_null_answers = True - self.do_dialogue = False - self.do_punctuation = False - self.channel_name = "Null" def is_expired(self, timeout_sec: float) -> bool: """ @@ -179,17 +152,14 @@ def to_dict(self) -> dict: dict с client_id, config, channel_name, audio_duration, флагами и накопленными результатами распознавания. """ - return { + data = super().to_dict() + data.update({ "client_id": self.client_id, - "config": self.config.model_dump() if self.config else None, - "channel_name": self.channel_name, - "audio_duration": self.audio_duration, - "do_dialogue": self.do_dialogue, - "do_punctuation": self.do_punctuation, "wait_null_answers": self.wait_null_answers, - "ws_collected_asr_res": self.ws_collected_asr_res, - "state": self.state.value, - } + "last_activity": self.last_activity, + "max_buffer_duration_sec": self.max_buffer_duration_sec, + }) + return data @classmethod def from_dict(cls, data: dict) -> "AudioSession": @@ -202,7 +172,7 @@ def from_dict(cls, data: dict) -> "AudioSession": Returns: AudioSession с восстановленной конфигурацией и флагами. """ - session = cls(client_id=data["client_id"]) + session = cls(client_id=data.get("client_id", data.get("session_id", ""))) if data.get("config"): session.config = WSConfigMessage(**data["config"]) session.channel_name = data.get("channel_name", "Null") @@ -210,6 +180,11 @@ def from_dict(cls, data: dict) -> "AudioSession": session.do_dialogue = data.get("do_dialogue", False) session.do_punctuation = data.get("do_punctuation", False) session.wait_null_answers = data.get("wait_null_answers", True) - session.ws_collected_asr_res = data.get("ws_collected_asr_res", {f"channel_{1}": []}) + session.collected_asr_res = data.get("collected_asr_res") or data.get("ws_collected_asr_res", {f"channel_{1}": []}) session.state = SessionState(data.get("state", "connecting")) + session.last_activity = data.get("last_activity", time.time()) + session.max_buffer_duration_sec = data.get( + "max_buffer_duration_sec", + getattr(settings, "WS_MAX_BUFFER_DURATION_SEC", 300.0), + ) return session diff --git a/utils/chunk_doing.py b/utils/chunk_doing.py index cbb1c96..37652e0 100644 --- a/utils/chunk_doing.py +++ b/utils/chunk_doing.py @@ -168,4 +168,112 @@ def samples_padding(samples, sample_rate = settings.BASE_SAMPLE_RATE, duration = # Идеальный размер padded_samples = samples - return padded_samples \ No newline at end of file + return padded_samples + + +async def find_last_speech_position_v2(session, is_last_chunk): + """ + Версия find_last_speech_position без глобальных dict. + Работает с полями session: audio_buffer, audio_overlap, audio_to_asr. + """ + if is_last_chunk: + last_audio = session.audio_overlap + session.audio_buffer + for i in range(0, len(last_audio), settings.MAX_OVERLAP_DURATION * 1000): + session.audio_to_asr.append( + last_audio[i:min(i + settings.MAX_OVERLAP_DURATION * 1000, len(last_audio))] + ) + else: + session.audio_buffer = session.audio_overlap + session.audio_buffer + frame_rate = session.audio_buffer.frame_rate + silero_bitrate = 16000 + + if not session.audio_buffer: + logger.error("Ошибка: audio_buffer пустой") + raise ValueError("audio_buffer не может быть пустым") + + if session.audio_buffer.frame_rate != silero_bitrate: + audio_for_vad = await async_resample_audiosegment(session.audio_buffer, silero_bitrate) + else: + audio_for_vad = session.audio_buffer + + logger.debug(f"Получено из буфера на обработку аудио продолжительностью {session.audio_buffer.duration_seconds}") + + audio_for_vad = session.audio_overlap + audio_for_vad + + try: + audio = get_np_array_samples_float32(audio_for_vad.raw_data, audio_for_vad.sample_width) + logger.debug(f"Аудио для VAD: длина={len(audio)}, min={np.min(audio)}, max={np.max(audio)}") + except Exception as e: + logger.error(f"Ошибка в get_np_array_samples_float32: {e}") + raise + + if np.any(np.isnan(audio)) or np.any(np.isinf(audio)): + logger.error("Обнаружены NaN или бесконечные значения в audio") + raise ValueError("Некорректные значения в audio") + + duration_seconds = 0.5 + frame_length = 512 if audio_for_vad.frame_rate == 16000 else 256 + + if frame_length is None: + raise ValueError("для VAD Поддерживаются только фреймрейты 8000 или 16000 Гц") + + frame_duration = frame_length / frame_rate + min_silence_frames = int(duration_seconds / frame_duration) + max_audio_length = len(audio) if len(audio) < settings.MAX_OVERLAP_DURATION * silero_bitrate else settings.MAX_OVERLAP_DURATION * silero_bitrate + partial_frame_length = 0 + + frames = [audio[i:i + frame_length] for i in range(int(len(audio) // 3), max_audio_length, frame_length)] + logger.debug(f"Создано фреймов: {len(frames)}, frame_length={frame_length}") + + silence_frames = 0 + await vad.reset_state() + vad_state = vad.state + + no_silent = False + for i, frame in enumerate(reversed(frames)): + vad.state = vad_state + try: + if len(frame) < frame_length: + partial_frame_length = len(frame) + logger.debug(f"Пропущен неполный фрейм: длина={partial_frame_length}") + continue + else: + logger.debug(f"Обработка фрейма {i}: длина={len(frame)}, min={np.min(frame)}, max={np.max(frame)}") + speech_prob, vad_state = await vad.is_speech(frame, audio_for_vad.frame_rate) + if speech_prob < vad.prob_level: + logger.debug(f"Найден не голос на speech_end = {max_audio_length-(i+1)*frame_length-partial_frame_length}") + silence_frames += 1 + if silence_frames >= min_silence_frames: + break + else: + silence_frames = 0 + logger.debug(f"Найден ГОЛОС на speech_end = {max_audio_length-i*frame_length-partial_frame_length}") + except Exception as e: + logger.error(f"Ошибка VAD - {e}" + f"\nframe_rate = {frame_rate}" + f"\nframe_length = {frame_length}" + f"\nframe_index = {i}" + f"\nframe_length_actual = {len(frame)}") + raise + else: + no_silent = True + + try: + if no_silent: + speech_end = max_audio_length + elif not partial_frame_length: + speech_end = max_audio_length - (i + 1) * frame_length + else: + speech_end = max_audio_length - i * frame_length + except Exception as e: + print(e) + else: + separation_time = int(speech_end * 1000 / silero_bitrate) + session.audio_to_asr.append(session.audio_buffer[:separation_time]) + session.audio_overlap = session.audio_buffer[separation_time:] + + logger.debug(f"Передано на ASR аудио продолжительностью {session.audio_to_asr[-1].duration_seconds}") + logger.debug(f"Передано в перекрытие аудио продолжительностью {session.audio_overlap.duration_seconds}") + session.audio_buffer = AudioSegment.silent(1, frame_rate) + + return diff --git a/utils/get_audio_file.py b/utils/get_audio_file.py index a4b8e23..3688cc1 100644 --- a/utils/get_audio_file.py +++ b/utils/get_audio_file.py @@ -7,26 +7,24 @@ from pydub import AudioSegment -from utils.pre_start_init import posted_and_downloaded_audio from utils.pre_start_init import paths -async def getting_audiofile(file_url, post_id) -> [bool, str]: +async def getting_audiofile(file_url, post_id) -> tuple[bool, str, io.BytesIO | None]: res = False error = str() file_ext = file_url.path.split('/')[-1].split('.')[-1] + buffer = None if file_ext in ['mp3', 'wav', 'ogg']: try: get_file_url = file_url.unicode_string() except Exception as e: logging.error(f"Error_url_parsing = {e}") - # get_file_url = file_url.geturl() + return False, f"Error_url_parsing = {e}", None else: with httpx.Client() as sess: try: - response = sess.get( - url=get_file_url - ) + response = sess.get(url=get_file_url) file_data = response.content except Exception as e: logging.error(f'Ошибка получения файла из ЕРП - {e}') @@ -34,29 +32,30 @@ async def getting_audiofile(file_url, post_id) -> [bool, str]: else: buffer = io.BytesIO(file_data) buffer.seek(0) - posted_and_downloaded_audio[post_id] = buffer res = True else: error = "No audio file in request link" - return res, error + return res, error, buffer # Todo - убрать корягу ниже -async def open_default_audiofile(post_id) -> tuple[bool, str]: +async def open_default_audiofile(post_id) -> tuple[bool, str, io.BytesIO | None]: res = False error_description = str() file = paths.get('test_file') file_ext = str(file).split('/')[-1].split('.')[-1] - buffer = io.BytesIO() + buffer = None if file_ext in ['mp3', 'wav', 'ogg']: try: - posted_and_downloaded_audio[post_id] = AudioSegment.from_file(file=file).export(buffer, format="wav") + buffer = io.BytesIO() + AudioSegment.from_file(file=file).export(buffer, format="wav") + buffer.seek(0) + res = True except Exception as e: logging.error(f"Error_file_opening = {e}") - else: - res = True + error_description = str(e) else: error_description = "No audio file in request link" - return res, error_description + return res, error_description, buffer From 54c579efe1f5bd0c1eef2344bbf598fcbfe53599 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Thu, 7 May 2026 17:27:32 +0300 Subject: [PATCH 29/46] =?UTF-8?q?=D0=B4=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=20=D1=80=D0=BE=D1=83=D1=82=D0=BE=D0=B2?= =?UTF-8?q?=20=D0=BC=D0=B5=D1=82=D1=80=D0=B8=D0=BA=20=D0=B4=D0=BB=D1=8F=20?= =?UTF-8?q?=D1=84=D1=80=D0=BE=D0=BD=D1=82=D0=B0.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/deps.py | 2 +- api/v1/endpoints/asr_ws.py | 10 +- api/v1/endpoints/root.py | 198 ++++++++++++++++++++++++++++++++++++- main.py | 4 +- models/ws_models.py | 3 + services/asr_pipeline.py | 13 +++ services/ws_manager.py | 53 ++++++++++ services/ws_metrics.py | 47 +++++++-- 8 files changed, 315 insertions(+), 15 deletions(-) diff --git a/api/deps.py b/api/deps.py index c4f1c60..c01908b 100644 --- a/api/deps.py +++ b/api/deps.py @@ -1,5 +1,5 @@ from datetime import datetime, timezone -from typing import Optional +from typing import Optional, Any from fastapi import Depends from fastapi.security import OAuth2PasswordBearer diff --git a/api/v1/endpoints/asr_ws.py b/api/v1/endpoints/asr_ws.py index a0dec23..6bc487b 100644 --- a/api/v1/endpoints/asr_ws.py +++ b/api/v1/endpoints/asr_ws.py @@ -110,7 +110,7 @@ async def websocket_endpoint( if message.get("bytes"): # Binary frame: отправляем сырые байты напрямую в pipeline, без base64-обёртки await process_audio_stream_chunk( - session, message["bytes"], recognizer, punctuator, manager + session, message["bytes"], recognizer, punctuator, manager, metrics ) continue elif message.get("text"): @@ -127,6 +127,10 @@ async def websocket_endpoint( ): await msg_router.route(msg, session, manager, metrics_collector=metrics) + # Подписка на периодический статус при запросе status + if msg.type == WSMessageType.status_request: + manager.set_subscribe_status(session.client_id, True) + # Копирование флагов из конфига в сессию (для ASR pipeline) if isinstance(msg, WSConfigMessage): session.wait_null_answers = msg.wait_null_answers @@ -141,12 +145,12 @@ async def websocket_endpoint( chunk_bytes = base64.b64decode(msg.audio_base64) if chunk_bytes: await process_audio_stream_chunk( - session, chunk_bytes, recognizer, punctuator, manager + session, chunk_bytes, recognizer, punctuator, manager, metrics ) # 7. Обработка конца потока (eos/eof) if isinstance(msg, WSEosMessage): - await process_final_audio(session, recognizer, punctuator, manager) + await process_final_audio(session, recognizer, punctuator, manager, metrics) break except WebSocketDisconnect: diff --git a/api/v1/endpoints/root.py b/api/v1/endpoints/root.py index abd9d0c..200b631 100644 --- a/api/v1/endpoints/root.py +++ b/api/v1/endpoints/root.py @@ -1,7 +1,9 @@ -from fastapi import APIRouter +from fastapi import APIRouter, Request +from fastapi.responses import HTMLResponse from models.fast_api_models import V1BaseResponse +from services.ws_metrics import SystemMetricsCollector from config import settings - +from static import mo router = APIRouter(prefix="", tags=["System"]) @@ -29,3 +31,195 @@ async def root(): "try_addr": f"http://{settings.HOST}:{settings.PORT}/docs" } ) + + +@router.get("/status", response_model=V1BaseResponse) +async def health_status(request: Request): + """ + REST endpoint для получения системных метрик (Задача 6.7). + Возвращает тот же JSON, что и WSStatusResponse. + """ + metrics: SystemMetricsCollector = request.app.state.metrics_collector + status = metrics.collect( + active_connections=request.app.state.ws_manager.active_connections_count, + max_connections=request.app.state.ws_manager.max_connections, + ) + return V1BaseResponse( + success=True, + error_description=None, + data=status.model_dump() + ) + + +@router.get("/monitor", response_class=HTMLResponse) +async def monitor_page(): + """ + Возвращает HTML-страницу мониторинга (Задача 6.7). + """ + html_content = """ + + + + + + ASR Monitor + + + +

ASR Monitor

+
+
+

GPU Memory

+
+
+
+
+

GPU Utilization

+
+
+
+
+

CPU Memory

+
+
+
+
+

CPU Utilization

+
+
+
+
+

Active Tasks

+
+
+
+

Connections

+
+
+
+

Adapter Status

+
idle
+
+
+

Uptime

+
+
+
+

Event Log

+
+ + + + + """ + return HTMLResponse(content=html_content) diff --git a/main.py b/main.py index 048c461..743fa0e 100644 --- a/main.py +++ b/main.py @@ -69,7 +69,7 @@ async def send_with_deprecation(message): if message["type"] == "http.response.start": headers = MutableHeaders(raw=message["headers"]) path = scope.get("path", "") - if path in {"/root/", "/root/demo", "/root/is_alive", "/root/post_file", "/root/post_one_step_req", "/root/ws"}: + if path in {"/", "/demo", "/is_alive", "/post_file", "/post_one_step_req", "/ws"}: headers["Deprecation"] = "true" headers["Warning"] = f'299 - "Legacy API is deprecated. Use /api/v1/ instead."' message["headers"] = headers.raw @@ -140,6 +140,7 @@ async def lifespan(app): # Graceful shutdown WebSocket (Задача 6.4, 6.9) if hasattr(app.state, "ws_manager"): + app.state.ws_manager.stop_status_broadcast() await app.state.ws_manager.disconnect_all() # cleanup (если нужно) @@ -155,7 +156,6 @@ async def lifespan(app): lifespan=lifespan, version="1.0", docs_url='/docs', - root_path='/root', title='ASR', description=WS_DESCRIPTION ) diff --git a/models/ws_models.py b/models/ws_models.py index a3cdb6e..79f824f 100644 --- a/models/ws_models.py +++ b/models/ws_models.py @@ -67,12 +67,15 @@ class WSStatusResponse(WSBaseMessage): adapter_status: Literal["idle", "busy", "overloaded"] = "idle" gpu_memory_free_mb: int | None = None gpu_memory_total_mb: int | None = None + gpu_utilization_percent: float | None = None cpu_memory_free_mb: int | None = None cpu_memory_total_mb: int | None = None + cpu_utilization_percent: float | None = None active_tasks_count: int = 0 active_connections_count: int = 0 queue_depth: int = 0 uptime_sec: float = 0.0 + temperature_celsius: float | None = None class WSWordItem(BaseModel): diff --git a/services/asr_pipeline.py b/services/asr_pipeline.py index d2fc602..5c03553 100644 --- a/services/asr_pipeline.py +++ b/services/asr_pipeline.py @@ -7,6 +7,7 @@ import logging from io import BytesIO +from typing import Union, Annotated, Optional, Any, List, Dict from pydub import AudioSegment @@ -34,6 +35,7 @@ async def process_audio_stream_chunk( recognizer, punctuator, manager: ConnectionManager, + metrics_collector: Optional[Any] = None, ) -> None: """ Обрабатывает входящий чанк аудио в потоковом режиме. @@ -56,6 +58,8 @@ async def process_audio_stream_chunk( punctuator: Экземпляр SbertPuncCaseOnnx (не используется в чанке, передаётся для единообразия). manager: Менеджер WebSocket-соединений для отправки ответов. """ + if metrics_collector is not None: + metrics_collector.increment_tasks() try: # --- 1. Проверка чётности --- if len(chunk_bytes) % 2 != 0: @@ -168,6 +172,9 @@ async def process_audio_stream_chunk( await manager.send_message(session.client_id, error_msg) except Exception: pass + finally: + if metrics_collector is not None: + metrics_collector.decrement_tasks() async def process_final_audio( @@ -175,6 +182,7 @@ async def process_final_audio( recognizer, punctuator, manager: ConnectionManager, + metrics_collector: Optional[Any] = None, ) -> None: """ Обрабатывает финальный буфер аудио по получении EOF/EOS. @@ -192,6 +200,8 @@ async def process_final_audio( punctuator: Экземпляр SbertPuncCaseOnnx. manager: Менеджер WebSocket-соединений. """ + if metrics_collector is not None: + metrics_collector.increment_tasks() try: # --- 1. Объединение остатков --- final_audio = session.audio_overlap + session.audio_buffer @@ -262,3 +272,6 @@ async def process_final_audio( await manager.send_message(session.client_id, error_msg) except Exception: pass + finally: + if metrics_collector is not None: + metrics_collector.decrement_tasks() diff --git a/services/ws_manager.py b/services/ws_manager.py index 01a3f52..4af8df6 100644 --- a/services/ws_manager.py +++ b/services/ws_manager.py @@ -63,6 +63,9 @@ def __init__(self, max_connections: int = 100) -> None: self.max_connections: int = max_connections self.active_connections: Dict[str, WebSocket] = {} self.connection_meta: Dict[str, ConnectionMeta] = {} + self._status_broadcast_task: Optional[asyncio.Task] = None + self._metrics_collector: Optional[Any] = None + self._broadcast_interval: float = 5.0 @property def active_connections_count(self) -> int: @@ -162,6 +165,56 @@ async def _send_text_safe(self, websocket: WebSocket, payload: str, client_id: s except Exception as exc: logger.debug("broadcast_status failed for %s: %s", client_id, exc) + def set_subscribe_status(self, client_id: str, value: bool) -> None: + """ + Устанавливает флаг подписки на периодический статус для клиента. + """ + meta = self.connection_meta.get(client_id) + if meta is not None: + meta.subscribe_status = value + + def start_status_broadcast( + self, + metrics_collector: Any, + interval_sec: float = 5.0, + ) -> None: + """ + Запускает фоновую задачу периодической рассылки статуса. + """ + self._metrics_collector = metrics_collector + self._broadcast_interval = interval_sec + if self._status_broadcast_task is None or self._status_broadcast_task.done(): + self._status_broadcast_task = asyncio.create_task( + self._status_broadcast_loop(), + name="ws_status_broadcast", + ) + logger.info("Started WS status broadcast every %.1f sec", interval_sec) + + def stop_status_broadcast(self) -> None: + """ + Останавливает фоновую задачу рассылки статуса. + """ + if self._status_broadcast_task is not None and not self._status_broadcast_task.done(): + self._status_broadcast_task.cancel() + logger.info("Stopped WS status broadcast") + + async def _status_broadcast_loop(self) -> None: + """Внутренний цикл периодической рассылки статуса.""" + try: + while True: + await asyncio.sleep(self._broadcast_interval) + if self._metrics_collector is None: + continue + status = self._metrics_collector.collect( + active_connections=self.active_connections_count, + max_connections=self.max_connections, + ) + await self.broadcast_status(status) + except asyncio.CancelledError: + logger.debug("Status broadcast loop cancelled") + except Exception as exc: + logger.error("Status broadcast loop error: %s", exc) + async def disconnect_all(self, code: int = 1001, reason: str = "Server shutdown") -> None: """ Принудительно закрывает все активные соединения. diff --git a/services/ws_metrics.py b/services/ws_metrics.py index 8a3222c..042e812 100644 --- a/services/ws_metrics.py +++ b/services/ws_metrics.py @@ -61,20 +61,40 @@ def get_gpu_stats(self) -> Tuple[Optional[int], Optional[int]]: logger.warning("Failed to get GPU stats: %s", exc) return None, None - def get_cpu_stats(self) -> Tuple[Optional[int], Optional[int]]: + def get_cpu_stats(self) -> Tuple[Optional[int], Optional[int], Optional[float]]: """ - Возвращает свободную и общую память RAM в мегабайтах. + Возвращает свободную и общую память RAM в мегабайтах и загрузку CPU в процентах. Returns: - Кортеж (free_mb, total_mb). Если psutil недоступен — (None, None). + Кортеж (free_mb, total_mb, cpu_percent). Если psutil недоступен — (None, None, None). """ try: import psutil mem = psutil.virtual_memory() free_mb = int(mem.available / 1024 / 1024) total_mb = int(mem.total / 1024 / 1024) - return free_mb, total_mb + cpu_percent = psutil.cpu_percent(interval=None) + return free_mb, total_mb, cpu_percent except Exception: + return None, None, None + + def get_gpu_utilization(self) -> Tuple[Optional[float], Optional[float]]: + """ + Возвращает загрузку GPU (utilization) и температуру в градусах Цельсия. + + Returns: + Кортеж (gpu_utilization_percent, temperature_celsius). Если pynvml недоступен — (None, None). + """ + if self._nvml_handle is None: + return None, None + try: + import pynvml + utilization = pynvml.nvmlDeviceGetUtilizationRates(self._nvml_handle) + gpu_util = float(utilization.gpu) + temperature = pynvml.nvmlDeviceGetTemperature(self._nvml_handle, pynvml.NVML_TEMPERATURE_GPU) + return gpu_util, float(temperature) + except Exception as exc: + logger.warning("Failed to get GPU utilization: %s", exc) return None, None def increment_tasks(self) -> None: @@ -108,6 +128,7 @@ def get_adapter_status( Пороги берутся из конфигурации: - GPU memory usage > WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT -> overloaded. + - GPU utilization > WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT -> overloaded. - active_connections / max_connections > WS_STATUS_BUSY_CONNECTIONS_THRESHOLD_PCT -> busy. Args: @@ -118,10 +139,18 @@ def get_adapter_status( Строка-статус: "idle", "busy" или "overloaded". """ gpu_free, gpu_total = self.get_gpu_stats() + gpu_util, _ = self.get_gpu_utilization() + + # Overloaded по памяти или utilization GPU if gpu_total and gpu_total > 0: gpu_used_pct = (gpu_total - gpu_free) / gpu_total * 100 - threshold = getattr(settings, "WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT", 90.0) - if gpu_used_pct >= threshold: + mem_threshold = getattr(settings, "WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT", 90.0) + if gpu_used_pct >= mem_threshold: + return "overloaded" + + if gpu_util is not None: + util_threshold = getattr(settings, "WS_STATUS_GPU_OVERLOAD_THRESHOLD_PCT", 90.0) + if gpu_util >= util_threshold: return "overloaded" if max_connections > 0: @@ -151,7 +180,8 @@ def collect( WSStatusResponse с актуальными метриками. """ gpu_free, gpu_total = self.get_gpu_stats() - cpu_free, cpu_total = self.get_cpu_stats() + gpu_util, gpu_temp = self.get_gpu_utilization() + cpu_free, cpu_total, cpu_util = self.get_cpu_stats() uptime = time.time() - self.start_time status = self.get_adapter_status(active_connections, max_connections) @@ -160,10 +190,13 @@ def collect( adapter_status=status, gpu_memory_free_mb=gpu_free, gpu_memory_total_mb=gpu_total, + gpu_utilization_percent=gpu_util, cpu_memory_free_mb=cpu_free, cpu_memory_total_mb=cpu_total, + cpu_utilization_percent=cpu_util, active_tasks_count=self._active_tasks, active_connections_count=active_connections, queue_depth=self.get_queue_depth(), uptime_sec=round(uptime, 2), + temperature_celsius=gpu_temp, ) From bf84292fcec7ee03d632c94df611b0c27ce390df Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Fri, 8 May 2026 11:01:05 +0300 Subject: [PATCH 30/46] =?UTF-8?q?=D0=A0=D0=B0=D0=B1=D0=BE=D1=82=D0=B0=20?= =?UTF-8?q?=D0=BD=D0=B0=D0=B4=20=D0=BE=D1=82=D0=BE=D0=B1=D1=80=D0=B0=D0=B6?= =?UTF-8?q?=D0=B5=D0=BD=D0=B8=D0=B5=D0=BC=20=D0=BC=D0=B5=D1=82=D1=80=D0=B8?= =?UTF-8?q?=D0=BA=20=D1=80=D0=B0=D0=B1=D0=BE=D1=82=D1=8B=20=D1=81=D0=B5?= =?UTF-8?q?=D1=80=D0=B2=D0=B8=D1=81=D0=B0.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Recognizer/engine/file_recognition.py | 2 +- Recognizer/engine/stream_recognition.py | 21 ++- api/v1/endpoints/root.py | 176 ------------------------ main.py | 6 + models/ws_models.py | 1 + services/asr_pipeline.py | 33 +++-- services/ws_metrics.py | 14 ++ 7 files changed, 57 insertions(+), 196 deletions(-) diff --git a/Recognizer/engine/file_recognition.py b/Recognizer/engine/file_recognition.py index ad4ec17..7104f3c 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -106,7 +106,7 @@ def process_file(session=None, recognizer=None, punctuator=None, diarizer=None, if params.use_batch: logger.info("Запрошен батчинг") - list_asr_result_wo_conf = simple_recognise_batch(session.audio_to_asr, params.batch_size, recognizer) # --> list + list_asr_result_wo_conf = asyncio.run(simple_recognise_batch(session.audio_to_asr, params.batch_size, recognizer)) # --> list for idx, asr_result_wo_conf in enumerate(list_asr_result_wo_conf): asr_result = recognizer.apply_postprocessing(asr_result_wo_conf, session.audio_duration) diff --git a/Recognizer/engine/stream_recognition.py b/Recognizer/engine/stream_recognition.py index 0e033ab..5b57f61 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -1,3 +1,4 @@ +import asyncio import time import numpy as np from config import settings @@ -41,8 +42,8 @@ def calc_speed(data): return speech_speed -async def simple_recognise(audio_data, recognizer) -> dict: - # Приводим фреймрейт к фреймрейту модели +def _simple_recognise_sync(audio_data, recognizer) -> dict: + """Синхронная реализация распознавания (CPU-bound, выполняется в отдельном потоке).""" if audio_data.frame_rate != settings.BASE_SAMPLE_RATE: audio_data = sync_resample_audiosegment(audio_data, settings.BASE_SAMPLE_RATE) @@ -53,6 +54,11 @@ async def simple_recognise(audio_data, recognizer) -> dict: return result +async def simple_recognise(audio_data, recognizer) -> dict: + """Асинхронная обёртка над CPU-bound распознаванием.""" + return await asyncio.to_thread(_simple_recognise_sync, audio_data, recognizer) + + async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_slow_down = False, recognizer = None) -> tuple: """ @@ -87,8 +93,8 @@ async def recognise_w_speed_correction(audio_data, multiplier=float(1.0), can_sl return result, speed, multiplier -def simple_recognise_batch(list_audio_data: list, batch_size: int = 8, recognizer=None) -> list: - +def _simple_recognise_batch_sync(list_audio_data: list, batch_size: int, recognizer) -> list: + """Синхронная реализация батч-распознавания (CPU-bound, выполняется в отдельном потоке).""" logger.info(f"Выполняется батчинг с размером {batch_size}") timer_sync_start = time.perf_counter() @@ -136,9 +142,12 @@ def simple_recognise_batch(list_audio_data: list, batch_size: int = 8, recognize logger.debug(f"Время на распознавание за {(time.perf_counter() - start_encoding):.4f} секунд.") + return list_of_dict_result - return list_of_dict_result +async def simple_recognise_batch(list_audio_data: list, batch_size: int = 8, recognizer=None) -> list: + """Асинхронная обёртка над CPU-bound батч-распознаванием.""" + return await asyncio.to_thread(_simple_recognise_batch_sync, list_audio_data, batch_size, recognizer) @@ -152,4 +161,4 @@ def simple_recognise_batch(list_audio_data: list, batch_size: int = 8, recognize "timestamps" : [ 0.64, 0.76, 0.92, 0.96, 1.08, 1.12, 1.24, 1.44, 1.48, 1.6, 1.64, 1.72, 1.76, 1.84, 1.92, 2.0, 2.2, 2.28, 2.32, 2.44, 2.52, 2.56, 2.68, 2.72, 2.92, 3.08, 3.16, 3.32, 3.36, 3.48, 3.6, 3.8, 3.84, 3.96, 4.0, 4.08, 4.32, 4.36, 4.48, 4.52, 4.6, 4.72, 4.8, 5.0, 5.08, 5.2, 5.24, 5.32, 5.68, 5.72, 5.84, 5.92, 5.96, 6.04, 6.16, 6.28, 6.32, 6.4, 6.52, 6.72, 6.8, 6.92, 7.0, 7.16, 7.2, 7.4, 7.56, 7.72, 7.8, 7.84, 7.96, 8.08, 8.24, 8.36, 8.4, 8.56, 8.64, 8.72, 8.8, 8.84, 8.96, 9.12, 9.16, 9.24, 9.4, 9.56, 9.6, 9.72, 9.84, 9.88, 9.96, 10.0, 10.12, 10.28, 10.36, 10.4, 10.52, 10.68, 10.76, 10.84, 10.96, 11.0, 11.16, 11.2, 11.36, 11.44, 11.76, 12.08, 12.16, 12.2, 12.32, 12.4, 12.44, 12.52, 12.68, 12.72, 12.8, 12.88, 12.92, 13.04, 13.12, 13.24, 13.36, 13.56, 13.6, 13.68, 13.76, 13.84, 13.88, 13.96, 14.04, 14.2, 14.28, 14.4, 14.44, 14.6, 14.64, 14.8, 14.96, 15.04, 15.08, 15.2, 15.4, 15.44, 15.6, 15.64, 15.76, 15.88, 15.96, 16.12, 16.24, 16.28, 16.56, 16.64, 16.72, 16.8, 16.92, 17.04, 17.08, 17.24, 17.28, 17.48, 17.72, 17.76, 17.92, 17.96, 18.08, 18.2, 18.32, 18.4, 18.44, 18.52 ], "tokens" : [ "я", " ", "б", "у", "д", "у", " ", "г", "о", "в", "о", "р", "и", "т", "ь", " ", "т", "р", "и", "д", "ц", "а", "т", "и", " ", "с", "е", "к", "у", "н", "д", "н", "ы", "м", "и", " ", "и", "н", "т", "е", "р", "в", "а", "л", "а", "м", "и", " ", "к", "а", "ж", "д", "ы", "й", " ", "р", "а", "з", " ", "п", "о", "в", "ы", "ш", "а", "я", " ", "с", "в", "о", "ю", " ", "с", "к", "о", "р", "о", "с", "т", "ь", " ", "н", "а", " ", "д", "в", "а", "д", "ц", "а", "т", "ь", " ", "с", "л", "о", "в", " ", "в", " ", "м", "и", "н", "у", "т", "у", " ", "п", "е", "р", "в", "ы", "й", " ", "и", "н", "т", "е", "р", "в", "а", "л", " ", "к", "о", "т", "о", "р", "ы", "й", " ", "к", "с", "т", "а", "т", "и", " ", "у", "ж", "е", " ", "н", "а", "ч", "а", "л", "с", "я", " ", "я", " ", "п", "р", "о", "и", "з", "н", "е", "с", "у", " ", "с", "о", "р", "о", "к", " ", "с", "л", "о", "в" ], "words" : [ ] - } \ No newline at end of file + } diff --git a/api/v1/endpoints/root.py b/api/v1/endpoints/root.py index 200b631..601e54a 100644 --- a/api/v1/endpoints/root.py +++ b/api/v1/endpoints/root.py @@ -1,9 +1,7 @@ from fastapi import APIRouter, Request -from fastapi.responses import HTMLResponse from models.fast_api_models import V1BaseResponse from services.ws_metrics import SystemMetricsCollector from config import settings -from static import mo router = APIRouter(prefix="", tags=["System"]) @@ -49,177 +47,3 @@ async def health_status(request: Request): error_description=None, data=status.model_dump() ) - - -@router.get("/monitor", response_class=HTMLResponse) -async def monitor_page(): - """ - Возвращает HTML-страницу мониторинга (Задача 6.7). - """ - html_content = """ - - - - - - ASR Monitor - - - -

ASR Monitor

-
-
-

GPU Memory

-
-
-
-
-

GPU Utilization

-
-
-
-
-

CPU Memory

-
-
-
-
-

CPU Utilization

-
-
-
-
-

Active Tasks

-
-
-
-

Connections

-
-
-
-

Adapter Status

-
idle
-
-
-

Uptime

-
-
-
-

Event Log

-
- - - - - """ - return HTMLResponse(content=html_content) diff --git a/main.py b/main.py index 743fa0e..b10667a 100644 --- a/main.py +++ b/main.py @@ -99,6 +99,12 @@ async def lifespan(app): app.state.state_store = InMemoryStateStore() logger.debug("WebSocket services initialized") + # Запуск фоновой push-рассылки статуса (Задача 6.6) + app.state.ws_manager.start_status_broadcast( + metrics_collector=app.state.metrics_collector, + interval_sec=settings.WS_STATUS_BROADCAST_INTERVAL_SEC, + ) + # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) diff --git a/models/ws_models.py b/models/ws_models.py index 79f824f..0bbcbeb 100644 --- a/models/ws_models.py +++ b/models/ws_models.py @@ -75,6 +75,7 @@ class WSStatusResponse(WSBaseMessage): active_connections_count: int = 0 queue_depth: int = 0 uptime_sec: float = 0.0 + uptime_formatted: str = "0s" temperature_celsius: float | None = None diff --git a/services/asr_pipeline.py b/services/asr_pipeline.py index 5c03553..fd8c10a 100644 --- a/services/asr_pipeline.py +++ b/services/asr_pipeline.py @@ -58,8 +58,6 @@ async def process_audio_stream_chunk( punctuator: Экземпляр SbertPuncCaseOnnx (не используется в чанке, передаётся для единообразия). manager: Менеджер WebSocket-соединений для отправки ответов. """ - if metrics_collector is not None: - metrics_collector.increment_tasks() try: # --- 1. Проверка чётности --- if len(chunk_bytes) % 2 != 0: @@ -121,12 +119,18 @@ async def process_audio_stream_chunk( if segment.duration_seconds <= 0: return - asr_result = await simple_recognise(segment, recognizer=recognizer) - asr_result_words = process_single_token_vocab_output(asr_result, session.audio_duration) - session.audio_duration += segment.duration_seconds + if metrics_collector is not None: + metrics_collector.increment_tasks() + try: + asr_result = await simple_recognise(segment, recognizer=recognizer) + asr_result_words = process_single_token_vocab_output(asr_result, session.audio_duration) + session.audio_duration += segment.duration_seconds - # Накопление для финального диалога - session.ws_collected_asr_res[f"channel_{1}"].append(asr_result_words) + # Накопление для финального диалога + session.ws_collected_asr_res[f"channel_{1}"].append(asr_result_words) + finally: + if metrics_collector is not None: + metrics_collector.decrement_tasks() # --- 7. Отправка результата --- text = asr_result_words.get("data", {}).get("text", "") @@ -172,9 +176,6 @@ async def process_audio_stream_chunk( await manager.send_message(session.client_id, error_msg) except Exception: pass - finally: - if metrics_collector is not None: - metrics_collector.decrement_tasks() async def process_final_audio( @@ -215,9 +216,15 @@ async def process_final_audio( logger.debug("Final audio padded with silence to %.3f sec", final_audio.duration_seconds) # --- 3. Распознавание --- - last_asr_result = await simple_recognise(final_audio, recognizer=recognizer) - last_result = process_single_token_vocab_output(last_asr_result, session.audio_duration) - session.ws_collected_asr_res[f"channel_{1}"].append(last_result) + if metrics_collector is not None: + metrics_collector.increment_tasks() + try: + last_asr_result = await simple_recognise(final_audio, recognizer=recognizer) + last_result = process_single_token_vocab_output(last_asr_result, session.audio_duration) + session.ws_collected_asr_res[f"channel_{1}"].append(last_result) + finally: + if metrics_collector is not None: + metrics_collector.decrement_tasks() # --- 4. Определение silence --- text = last_result.get("data", {}).get("text", "") diff --git a/services/ws_metrics.py b/services/ws_metrics.py index 042e812..e0e9f4c 100644 --- a/services/ws_metrics.py +++ b/services/ws_metrics.py @@ -164,6 +164,19 @@ def get_adapter_status( return "idle" + @staticmethod + def _format_uptime(seconds: float) -> str: + total = int(seconds) + hours, rem = divmod(total, 3600) + minutes, secs = divmod(rem, 60) + parts = [] + if hours > 0: + parts.append(f"{hours}h") + if minutes > 0 or hours > 0: + parts.append(f"{minutes:02d}m") + parts.append(f"{secs:02d}s") + return " ".join(parts) + def collect( self, active_connections: int = 0, @@ -198,5 +211,6 @@ def collect( active_connections_count=active_connections, queue_depth=self.get_queue_depth(), uptime_sec=round(uptime, 2), + uptime_formatted=self._format_uptime(uptime), temperature_celsius=gpu_temp, ) From 488817f0a32b2f3c025dfb180950f6a866253dff Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Fri, 8 May 2026 11:12:43 +0300 Subject: [PATCH 31/46] =?UTF-8?q?=D0=A7=D0=B8=D1=81=D1=82=D0=BA=D0=B0.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- models/fast_api_models.py | 82 --------------------------------------- 1 file changed, 82 deletions(-) diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 3ef5df1..df5deac 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -172,85 +172,3 @@ class PostFileRequest(BaseModel): } } ) - -class WebSocketModel(BaseModel): - """OpenAPI не хочет описывать WS, а я не хочу изучать OPEN API. По этому описание тут. - - Подключение на порт: 49153 - На вход жду поток binary, buffer_size +- 6400, mono, wav. - - Протокол обмена сообщениями: - 1. Начальное сообщение (конфигурация): - {'text': { "config" : { "sample_rate" : any(int/float), "wait_null_answers": Bool, - "do_dialogue": Bool, "do_punctuation": Bool}}} - do_punctuation отработает только если do_dialogue = True - - 2. Последующие сообщения с аудио-данными: - {"bytes": binary} - - 3. Последнее сообщение (сигнал окончания передачи): - {'text': '{ "eof" : 1}'} - - Формат ответа от сервера: - {"silence": Bool, "data": str, "error": None/str, "last_message": Bool, - "sentenced_data": {}} - - Пример ответа "data": { - "result" : [{ - "conf" : 1.000000, - "end" : 3.120000, - "start" : 2.340000, - "word" : "здравствуйте" - }, { - "conf" : 1.000000, - "end" : 3.870000, - "start" : 3.600000, - "word" : "вы" - }, - ... - { - "conf" : 0.994019, - "end" : 11.790000, - "start" : 10.890000, - "word" : "записываются" - }], - "text" : "здравствуйте вы ... записываются" - } - - Пример ответа "sentenced_data": { - "raw_text_sentenced_recognition": "channel_1: Татьяна, добрый день. Меня зовут Ульяна.'/n'channel_1: Звоню уточнить по поводу документов.", - "list_of_sentenced_recognitions": [ - { - "start": 2.28, - "text": "Татьяна, добрый день. Меня зовут Ульяна.", - "speaker": "channel_1" - }, - { - "start": 8.24, - "text": "Звоню уточнить по поводу документов.", - "speaker": "channel_1" - }, - ], - "full_text_only": [ - "Татьяна, добрый день. Меня зовут Ульяна. Звоню уточнить по поводу документов." - ], - "err_state": null - } - """ - - model_config = ConfigDict( - json_schema_extra={ - "example": { - "text": { - "config": { - "sample_rate": 8000, - "wait_null_answers": False, - "do_dialogue": True, - "do_punctuation": True - } - } - } - } - ) - - pass From 416d152410f919ff4b93aa942afa8e711ed935fd Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Sat, 9 May 2026 21:58:39 +0300 Subject: [PATCH 32/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B0=20web-=D0=BF=D0=B0=D0=BD=D0=B5=D0=BB=D1=8C.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- alembic.ini | 108 +++++ alembic/README | 1 + alembic/__init__.py | 0 alembic/env.py | 82 ++++ alembic/script.py.mako | 26 ++ api/v1/api.py | 6 + api/v1/endpoints/admin.py | 557 +++++++++++++++++++++++++ api/v1/endpoints/admin_ws.py | 66 +++ api/v1/endpoints/asr_file.py | 32 +- api/v1/endpoints/asr_url.py | 32 +- api/v1/endpoints/asr_ws.py | 105 ++++- api/v1/endpoints/auth.py | 205 ++++++++++ api/v1/endpoints/tg.py | 54 +++ api/v1/endpoints/user.py | 373 +++++++++++++++++ core/deps.py | 144 +++++++ core/middleware.py | 84 ++++ db/__init__.py | 1 + db/base.py | 24 ++ db/enums.py | 51 +++ db/models.py | 228 +++++++++++ db/session.py | 40 ++ main.py | 31 ++ models/admin.py | 162 ++++++++ models/auth.py | 38 ++ models/user.py | 98 +++++ routes/admin.py | 108 +++++ routes/user.py | 84 ++++ services/admin_service.py | 329 +++++++++++++++ services/asr_pipeline.py | 48 ++- services/auth_service.py | 146 +++++++ services/metrics_reporter.py | 50 +++ services/ws_manager.py | 1 + services/ws_session.py | 1 + static/css/design-system.css | 147 +++++++ static/js/asr_client.js | 113 ++++++ static/js/asr_page.js | 296 ++++++++++++++ static/js/auth.js | 89 ++++ static/js/ui.js | 83 ++++ static/monitor.html | 291 +++++++++++++ templates/admin/api_keys.html | 32 ++ templates/admin/base_admin.html | 35 ++ templates/admin/dashboard.html | 18 + templates/admin/login.html | 41 ++ templates/admin/logs.html | 31 ++ templates/admin/sessions.html | 8 + templates/admin/settings.html | 21 + templates/admin/subscriptions.html | 32 ++ templates/admin/tariffs.html | 32 ++ templates/admin/telegram.html | 21 + templates/admin/transactions.html | 33 ++ templates/admin/users.html | 34 ++ templates/auth/login.html | 66 +++ templates/auth/register.html | 72 ++++ templates/base.html | 78 ++++ templates/tg/index.html | 41 ++ templates/user/api_keys.html | 155 +++++++ templates/user/asr.html | 145 +++++++ templates/user/base_user.html | 44 ++ templates/user/dashboard.html | 132 ++++++ templates/user/history.html | 156 +++++++ templates/user/history_detail.html | 104 +++++ templates/user/profile.html | 131 ++++++ templates/user/subscription.html | 91 +++++ to_do_front.txt | 631 +++++++++++++++++++++++++++++ utils/chunk_doing.py | 204 ++++++---- 65 files changed, 6626 insertions(+), 96 deletions(-) create mode 100644 alembic.ini create mode 100644 alembic/README create mode 100644 alembic/__init__.py create mode 100644 alembic/env.py create mode 100644 alembic/script.py.mako create mode 100644 api/v1/endpoints/admin.py create mode 100644 api/v1/endpoints/admin_ws.py create mode 100644 api/v1/endpoints/auth.py create mode 100644 api/v1/endpoints/tg.py create mode 100644 api/v1/endpoints/user.py create mode 100644 core/deps.py create mode 100644 core/middleware.py create mode 100644 db/__init__.py create mode 100644 db/base.py create mode 100644 db/enums.py create mode 100644 db/models.py create mode 100644 db/session.py create mode 100644 models/admin.py create mode 100644 models/auth.py create mode 100644 models/user.py create mode 100644 routes/admin.py create mode 100644 routes/user.py create mode 100644 services/admin_service.py create mode 100644 services/auth_service.py create mode 100644 services/metrics_reporter.py create mode 100644 static/css/design-system.css create mode 100644 static/js/asr_client.js create mode 100644 static/js/asr_page.js create mode 100644 static/js/auth.js create mode 100644 static/js/ui.js create mode 100644 static/monitor.html create mode 100644 templates/admin/api_keys.html create mode 100644 templates/admin/base_admin.html create mode 100644 templates/admin/dashboard.html create mode 100644 templates/admin/login.html create mode 100644 templates/admin/logs.html create mode 100644 templates/admin/sessions.html create mode 100644 templates/admin/settings.html create mode 100644 templates/admin/subscriptions.html create mode 100644 templates/admin/tariffs.html create mode 100644 templates/admin/telegram.html create mode 100644 templates/admin/transactions.html create mode 100644 templates/admin/users.html create mode 100644 templates/auth/login.html create mode 100644 templates/auth/register.html create mode 100644 templates/base.html create mode 100644 templates/tg/index.html create mode 100644 templates/user/api_keys.html create mode 100644 templates/user/asr.html create mode 100644 templates/user/base_user.html create mode 100644 templates/user/dashboard.html create mode 100644 templates/user/history.html create mode 100644 templates/user/history_detail.html create mode 100644 templates/user/profile.html create mode 100644 templates/user/subscription.html create mode 100644 to_do_front.txt diff --git a/alembic.ini b/alembic.ini new file mode 100644 index 0000000..8f6234e --- /dev/null +++ b/alembic.ini @@ -0,0 +1,108 @@ +# A generic, single database configuration. + +[alembic] +# path to migration scripts +script_location = alembic + +# template used to generate migration file names; The default value is %%(rev)s_%%(slug)s +# file_template = %%(rev)s_%%(slug)s + +# sys.path path, will be prepended to sys.path if present. +# defaults to the current working directory. +prepend_sys_path = . + +# timezone to use when rendering the date within the migration file +# as well as the filename. +# string value is passed to dateutil.tz.gettz() +# leave blank for localtime +# timezone = + +# max length of characters to apply to the +# "slug" field +# truncate_slug_length = 40 + +# set to 'true' to run the environment during +# the 'revision' command, regardless of autogenerate +# revision_environment = false + +# set to 'true' to allow .pyc and .pyo files without +# a source .py file to be detected as revisions in the +# versions/ directory +# sourceless = false + +# version path separator; As mentioned above, this is the character used to split +# version_locations. The default within new alembic.ini files is "os", which uses +# os.pathsep. If this key is omitted entirely, it falls back to the legacy +# behaviour of splitting on spaces and/or commas. +# Valid values for version_path_separator are: +# +# version_path_separator = : +# version_path_separator = ; +# version_path_separator = space +version_path_separator = os + +# set to 'true' to search source files recursively +# in each "version_locations" directory +# new in Alembic version 1.10 +# recursive_version_locations = false + +# the output encoding used when revision files +# are written from script.py.mako +# output_encoding = utf-8 + +# Для локальной разработки — SQLite (async). В production переопределить +# через переменную окружения DATABASE_URL (postgresql+asyncpg://...). +sqlalchemy.url = sqlite+aiosqlite:///./asr_local.db + + +[post_write_hooks] +# post_write_hooks defines scripts or Python functions that are run +# on newly generated revision scripts. See the documentation for further +# detail and examples + +# format using "black" - use the console_scripts runner, against the "black" entrypoint +# hooks = black +# black.type = console_scripts +# black.entrypoint = black +# black.options = -l 79 REVISION_SCRIPT_FILENAME + +# lint with attempts to fix using "ruff" - use the exec runner, execute a binary +# hooks = ruff +# ruff.type = exec +# ruff.executable = %(here)s/.venv/bin/ruff +# ruff.options = --fix REVISION_SCRIPT_FILENAME + +# Logging configuration +[loggers] +keys = root,sqlalchemy,alembic + +[handlers] +keys = console + +[formatters] +keys = generic + +[logger_root] +level = WARN +handlers = console +qualname = + +[logger_sqlalchemy] +level = WARN +handlers = +qualname = sqlalchemy.engine + +[logger_alembic] +level = INFO +handlers = +qualname = alembic + +[handler_console] +class = StreamHandler +args = (sys.stderr,) +level = NOTSET +formatter = generic + +[formatter_generic] +format = %(levelname)-5.5s [%(name)s] %(message)s +datefmt = %H:%M:%S diff --git a/alembic/README b/alembic/README new file mode 100644 index 0000000..4a8bfb4 --- /dev/null +++ b/alembic/README @@ -0,0 +1 @@ +Generic single-database configuration with async SQLAlchemy 2.0. diff --git a/alembic/__init__.py b/alembic/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/alembic/env.py b/alembic/env.py new file mode 100644 index 0000000..75d2c2e --- /dev/null +++ b/alembic/env.py @@ -0,0 +1,82 @@ +"""Alembic environment (async).""" + +import asyncio +from logging.config import fileConfig + +from sqlalchemy import pool +from sqlalchemy.engine import Connection +from sqlalchemy.ext.asyncio import async_engine_from_config + +from alembic import context + +import os + +# Импорт базового класса и регистрация моделей в metadata +from db.base import Base +from db import models # noqa: F401 — регистрирует ORM-модели + +# this is the Alembic Config object, which provides +# access to the values within the .ini file in use. +config = context.config + +# Interpret the config file for Python logging. +# This line sets up loggers basically. +if config.config_file_name is not None: + fileConfig(config.config_file_name) + +# Если задана переменная окружения DATABASE_URL — переопределяем URL из alembic.ini. +# Иначе используем то, что указано в alembic.ini (например, SQLite для локальной разработки). +DATABASE_URL = os.getenv("DATABASE_URL") +if DATABASE_URL: + config.set_main_option("sqlalchemy.url", DATABASE_URL) + +# add your model's MetaData object here +# for 'autogenerate' support +target_metadata = Base.metadata + + +def run_migrations_offline() -> None: + """Run migrations in 'offline' mode.""" + url = config.get_main_option("sqlalchemy.url") + context.configure( + url=url, + target_metadata=target_metadata, + literal_binds=True, + dialect_opts={"paramstyle": "named"}, + ) + + with context.begin_transaction(): + context.run_migrations() + + +def do_run_migrations(connection: Connection) -> None: + """Wrapper для синхронного вызова внутри async соединения.""" + context.configure(connection=connection, target_metadata=target_metadata) + + with context.begin_transaction(): + context.run_migrations() + + +async def run_async_migrations() -> None: + """Run migrations in 'online' mode with async engine.""" + connectable = async_engine_from_config( + config.get_section(config.config_ini_section, {}), + prefix="sqlalchemy.", + poolclass=pool.NullPool, + ) + + async with connectable.connect() as connection: + await connection.run_sync(do_run_migrations) + + await connectable.dispose() + + +def run_migrations_online() -> None: + """Entrypoint для online-миграций.""" + asyncio.run(run_async_migrations()) + + +if context.is_offline_mode(): + run_migrations_offline() +else: + run_migrations_online() diff --git a/alembic/script.py.mako b/alembic/script.py.mako new file mode 100644 index 0000000..fbc4b07 --- /dev/null +++ b/alembic/script.py.mako @@ -0,0 +1,26 @@ +"""${message} + +Revision ID: ${up_revision} +Revises: ${down_revision | comma,n} +Create Date: ${create_date} + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +${imports if imports else ""} + +# revision identifiers, used by Alembic. +revision: str = ${repr(up_revision)} +down_revision: Union[str, None] = ${repr(down_revision)} +branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} +depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} + + +def upgrade() -> None: + ${upgrades if upgrades else "pass"} + + +def downgrade() -> None: + ${downgrades if downgrades else "pass"} diff --git a/api/v1/api.py b/api/v1/api.py index 0d864ed..a3fc1ef 100644 --- a/api/v1/api.py +++ b/api/v1/api.py @@ -5,6 +5,9 @@ from api.v1.endpoints.asr_file import router as asr_file_router from api.v1.endpoints.asr_ws import router as asr_ws_router from api.v1.endpoints.health import router as health_router +from api.v1.endpoints.auth import router as auth_router +from api.v1.endpoints.user import router as user_router +from api.v1.endpoints.admin import router as admin_api_router router = APIRouter(prefix="/api/v1") @@ -13,3 +16,6 @@ router.include_router(asr_file_router) router.include_router(asr_ws_router) router.include_router(health_router) +router.include_router(auth_router) +router.include_router(user_router) +router.include_router(admin_api_router) diff --git a/api/v1/endpoints/admin.py b/api/v1/endpoints/admin.py new file mode 100644 index 0000000..6ba2bf9 --- /dev/null +++ b/api/v1/endpoints/admin.py @@ -0,0 +1,557 @@ +"""FastAPI-роутер админ-панели.""" + +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, status +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from core.deps import get_current_user, require_admin, require_superadmin +from db.models import ApiKey, Plan, Subscription, Transaction, User +from db.session import get_db_session +from models.admin import ( + AdminApiKeyResponse, + AdminAuditLogResponse, + AdminBroadcastRequest, + AdminMaintenanceToggle, + AdminMetricsResponse, + AdminPlanCreateUpdateRequest, + AdminPlanResponse, + AdminSubscriptionResponse, + AdminSystemLogResponse, + AdminTelegramConfig, + AdminTelegramStats, + AdminTransactionResponse, + AdminUserDetailResponse, + AdminUserListItem, + AdminUserUpdateRequest, + PaginationParams, +) +from services import admin_service + +router = APIRouter(prefix="/admin", tags=["admin"]) + + +@router.get("/metrics", response_model=AdminMetricsResponse) +async def admin_metrics(current_user: User = Depends(require_admin)): + """Текущие метрики системы (заглушка).""" + return AdminMetricsResponse() + + +@router.get("/metrics/history") +async def admin_metrics_history(current_user: User = Depends(require_admin)): + """История метрик за период (заглушка).""" + return {"detail": "История метрик — заглушка"} + + +@router.get("/users", response_model=list[AdminUserListItem]) +async def admin_users_list( + pagination: PaginationParams = Depends(), + search: Optional[str] = None, + role: Optional[str] = None, + is_active: Optional[bool] = None, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Список пользователей с пагинацией и фильтрами.""" + users, total = await admin_service.get_users_list( + db, + page=pagination.page, + per_page=pagination.per_page, + search=search, + role=role, + is_active=is_active, + ) + return [ + AdminUserListItem( + id=u.id, + email=u.email, + full_name=u.full_name, + role=u.role.value, + is_active=u.is_active, + created_at=u.created_at, + last_login_at=u.last_login_at, + telegram_linked=u.telegram_id is not None, + ) + for u in users + ] + + +@router.get("/users/{user_id}", response_model=AdminUserDetailResponse) +async def admin_user_detail( + user_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Детали пользователя.""" + user = await admin_service.get_user_by_id(db, user_id) + if not user: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Пользователь не найден" + ) + return AdminUserDetailResponse( + id=user.id, + email=user.email, + full_name=user.full_name, + role=user.role.value, + is_active=user.is_active, + created_at=user.created_at, + last_login_at=user.last_login_at, + telegram_linked=user.telegram_id is not None, + phone=user.phone, + telegram_id=user.telegram_id, + telegram_username=user.telegram_username, + ) + + +@router.put("/users/{user_id}", response_model=AdminUserDetailResponse) +async def admin_user_update( + user_id: str, + payload: AdminUserUpdateRequest, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Редактирование пользователя админом.""" + user = await admin_service.get_user_by_id(db, user_id) + if not user: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Пользователь не найден" + ) + user = await admin_service.update_user( + db, user, payload.model_dump(exclude_unset=True) + ) + return AdminUserDetailResponse( + id=user.id, + email=user.email, + full_name=user.full_name, + role=user.role.value, + is_active=user.is_active, + created_at=user.created_at, + last_login_at=user.last_login_at, + telegram_linked=user.telegram_id is not None, + phone=user.phone, + telegram_id=user.telegram_id, + telegram_username=user.telegram_username, + ) + + +@router.delete("/users/{user_id}") +async def admin_user_delete( + user_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Soft-delete / блокировка пользователя.""" + user = await admin_service.get_user_by_id(db, user_id) + if not user: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Пользователь не найден" + ) + user.is_active = False + await db.commit() + return {"detail": "Пользователь деактивирован"} + + +@router.get("/users/{user_id}/sessions") +async def admin_user_sessions( + user_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), + limit: int = 50, +): + """История ASR-сессий пользователя.""" + sessions = await admin_service.get_user_sessions(db, user_id, limit=limit) + return [ + { + "id": s.id, + "session_type": s.session_type.value, + "status": s.status.value, + "audio_duration_sec": s.audio_duration_sec, + "created_at": s.created_at, + "completed_at": s.completed_at, + } + for s in sessions + ] + + +@router.post("/users/{user_id}/impersonate") +async def admin_user_impersonate( + user_id: str, + current_user: User = Depends(require_superadmin), + db: AsyncSession = Depends(get_db_session), +): + """Получить access token от имени пользователя (superadmin only).""" + user = await admin_service.get_user_by_id(db, user_id) + if not user: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Пользователь не найден" + ) + token = await admin_service.impersonate_user(user) + return {"access_token": token, "token_type": "bearer"} + + +@router.get("/tariffs", response_model=list[AdminPlanResponse]) +async def admin_tariffs_list( + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Список тарифных планов.""" + plans = await admin_service.get_plans(db) + return [ + AdminPlanResponse( + id=p.id, + code=p.code, + name=p.name, + description=p.description, + max_requests_per_minute=p.max_requests_per_minute, + max_audio_duration_sec=p.max_audio_duration_sec, + price_per_month=float(p.price_per_month) + if p.price_per_month is not None + else None, + is_active=p.is_active, + created_at=p.created_at, + updated_at=p.updated_at, + ) + for p in plans + ] + + +@router.post("/tariffs", response_model=AdminPlanResponse) +async def admin_tariff_create( + payload: AdminPlanCreateUpdateRequest, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Создать тарифный план.""" + plan = await admin_service.create_plan(db, payload.model_dump()) + return AdminPlanResponse( + id=plan.id, + code=plan.code, + name=plan.name, + description=plan.description, + max_requests_per_minute=plan.max_requests_per_minute, + max_audio_duration_sec=plan.max_audio_duration_sec, + price_per_month=float(plan.price_per_month) + if plan.price_per_month is not None + else None, + is_active=plan.is_active, + created_at=plan.created_at, + updated_at=plan.updated_at, + ) + + +@router.put("/tariffs/{plan_id}", response_model=AdminPlanResponse) +async def admin_tariff_update( + plan_id: str, + payload: AdminPlanCreateUpdateRequest, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Обновить тарифный план.""" + result = await db.execute(select(Plan).where(Plan.id == plan_id)) + plan = result.scalar_one_or_none() + if not plan: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Тариф не найден" + ) + plan = await admin_service.update_plan( + db, plan, payload.model_dump(exclude_unset=True) + ) + return AdminPlanResponse( + id=plan.id, + code=plan.code, + name=plan.name, + description=plan.description, + max_requests_per_minute=plan.max_requests_per_minute, + max_audio_duration_sec=plan.max_audio_duration_sec, + price_per_month=float(plan.price_per_month) + if plan.price_per_month is not None + else None, + is_active=plan.is_active, + created_at=plan.created_at, + updated_at=plan.updated_at, + ) + + +@router.delete("/tariffs/{plan_id}") +async def admin_tariff_delete( + plan_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Деактивировать тариф.""" + result = await db.execute(select(Plan).where(Plan.id == plan_id)) + plan = result.scalar_one_or_none() + if not plan: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Тариф не найден" + ) + await admin_service.delete_plan(db, plan) + return {"detail": "Тариф деактивирован"} + + +@router.get("/subscriptions", response_model=list[AdminSubscriptionResponse]) +async def admin_subscriptions_list( + pagination: PaginationParams = Depends(), + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Список подписок.""" + subs, total = await admin_service.get_subscriptions( + db, page=pagination.page, per_page=pagination.per_page + ) + return [ + AdminSubscriptionResponse( + id=s.id, + user_id=s.user_id, + plan_id=s.plan_id, + plan_name=None, + status=s.status.value, + started_at=s.started_at, + expires_at=s.expires_at, + auto_renew=s.auto_renew, + ) + for s in subs + ] + + +@router.post("/subscriptions/{sub_id}/extend") +async def admin_subscription_extend( + sub_id: str, + days: int = 30, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Ручное продление подписки.""" + result = await db.execute(select(Subscription).where(Subscription.id == sub_id)) + sub = result.scalar_one_or_none() + if not sub: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Подписка не найдена" + ) + sub = await admin_service.extend_subscription(db, sub, days=days) + return {"detail": f"Подписка продлена на {days} дней", "expires_at": sub.expires_at} + + +@router.post("/subscriptions/{sub_id}/cancel") +async def admin_subscription_cancel( + sub_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Ручная отмена подписки.""" + result = await db.execute(select(Subscription).where(Subscription.id == sub_id)) + sub = result.scalar_one_or_none() + if not sub: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Подписка не найдена" + ) + await admin_service.cancel_subscription_admin(db, sub) + return {"detail": "Подписка отменена"} + + +@router.get("/transactions", response_model=list[AdminTransactionResponse]) +async def admin_transactions_list( + pagination: PaginationParams = Depends(), + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Список транзакций.""" + txs, total = await admin_service.get_transactions( + db, page=pagination.page, per_page=pagination.per_page + ) + return [ + AdminTransactionResponse( + id=t.id, + user_id=t.user_id, + subscription_id=t.subscription_id, + amount=float(t.amount) if t.amount is not None else None, + currency=t.currency, + status=t.status.value, + payment_provider=t.payment_provider, + external_payment_id=t.external_payment_id, + created_at=t.created_at, + ) + for t in txs + ] + + +@router.get("/logs", response_model=list[AdminSystemLogResponse]) +async def admin_logs_list( + pagination: PaginationParams = Depends(), + level: Optional[str] = None, + component: Optional[str] = None, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Системные логи.""" + logs, total = await admin_service.get_system_logs( + db, + page=pagination.page, + per_page=pagination.per_page, + level=level, + component=component, + ) + return [ + AdminSystemLogResponse( + id=l.id, + level=l.level.value, + component=l.component, + message=l.message, + meta=l.meta, + created_at=l.created_at, + ) + for l in logs + ] + + +@router.get("/audit", response_model=list[AdminAuditLogResponse]) +async def admin_audit_list( + pagination: PaginationParams = Depends(), + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Аудит действий админов.""" + logs, total = await admin_service.get_audit_logs( + db, page=pagination.page, per_page=pagination.per_page + ) + return [ + AdminAuditLogResponse( + id=l.id, + admin_id=l.admin_id, + action=l.action, + target_type=l.target_type, + target_id=l.target_id, + details=l.details, + created_at=l.created_at, + ) + for l in logs + ] + + +@router.get("/api-keys", response_model=list[AdminApiKeyResponse]) +async def admin_api_keys_list( + pagination: PaginationParams = Depends(), + user_id: Optional[str] = None, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Все API-ключи (с фильтром по пользователю).""" + keys, total = await admin_service.get_api_keys( + db, page=pagination.page, per_page=pagination.per_page, user_id=user_id + ) + return [ + AdminApiKeyResponse( + id=k.id, + user_id=k.user_id, + user_email=None, + name=k.name, + is_active=k.is_active, + created_at=k.created_at, + last_used_at=k.last_used_at, + ) + for k in keys + ] + + +@router.delete("/api-keys/{key_id}") +async def admin_api_key_revoke( + key_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Отозвать API-ключ админом.""" + result = await db.execute(select(ApiKey).where(ApiKey.id == key_id)) + key = result.scalar_one_or_none() + if not key: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Ключ не найден" + ) + await admin_service.revoke_api_key_admin(db, key) + return {"detail": "Ключ отозван"} + + +@router.get("/queue") +async def admin_queue(current_user: User = Depends(require_admin)): + """Текущая очередь задач.""" + return await admin_service.get_queue_status() + + +@router.post("/queue/{task_id}/cancel") +async def admin_queue_cancel( + task_id: str, + current_user: User = Depends(require_admin), +): + """Отменить задачу.""" + success = await admin_service.cancel_task(task_id) + return { + "detail": "Задача отменена" if success else "Не удалось отменить задачу" + } + + +@router.post("/users/{user_id}/sessions/{session_id}/disconnect") +async def admin_disconnect_session( + user_id: str, + session_id: str, + current_user: User = Depends(require_admin), +): + """Принудительно закрыть WS-сессию пользователя.""" + success = await admin_service.disconnect_user_session(user_id, session_id) + return { + "detail": "Сессия закрыта" if success else "Не удалось закрыть сессию" + } + + +@router.post("/maintenance") +async def admin_maintenance( + payload: AdminMaintenanceToggle, + current_user: User = Depends(require_admin), +): + """Включить/выключить режим обслуживания.""" + admin_service.set_maintenance_mode(payload.enabled) + return { + "detail": f"Режим обслуживания {'включён' if payload.enabled else 'выключён'}" + } + + +@router.get("/telegram/config", response_model=AdminTelegramConfig) +async def admin_telegram_config( + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Настройки Telegram-бота.""" + cfg = await admin_service.get_telegram_config(db) + if not cfg: + return AdminTelegramConfig() + return AdminTelegramConfig(webapp_url=cfg.webapp_url, is_active=cfg.is_active) + + +@router.post("/telegram/webhook") +async def admin_telegram_webhook( + url: str, + current_user: User = Depends(require_admin), +): + """Установить webhook бота (заглушка).""" + import os + + bot_token = os.getenv("TELEGRAM_BOT_TOKEN", "") + return await admin_service.set_telegram_webhook(url, bot_token) + + +@router.get("/telegram/stats", response_model=AdminTelegramStats) +async def admin_telegram_stats( + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Статистика Telegram Web App.""" + stats = await admin_service.get_telegram_stats(db) + return AdminTelegramStats(**stats) + + +@router.post("/telegram/broadcast") +async def admin_telegram_broadcast( + payload: AdminBroadcastRequest, + current_user: User = Depends(require_admin), +): + """Рассылка сообщения всем пользователям бота (заглушка).""" + return await admin_service.broadcast_message(payload.message) diff --git a/api/v1/endpoints/admin_ws.py b/api/v1/endpoints/admin_ws.py new file mode 100644 index 0000000..69e015f --- /dev/null +++ b/api/v1/endpoints/admin_ws.py @@ -0,0 +1,66 @@ +"""WebSocket endpoint для real-time метрик админ-панели.""" + +import asyncio +import logging + +from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status +from sqlalchemy import select + +from core.security import decode_token # type: ignore[import-untyped] +from db.enums import UserRole +from db.models import User +from db.session import AsyncSessionLocal + +router = APIRouter() + + +@router.websocket("/api/v1/admin/ws") +async def admin_websocket(websocket: WebSocket): + """Push метрик для админов. Auth через первый фрейм.""" + await websocket.accept() + + # Ожидаем auth фрейм в течение 5 секунд + try: + auth_msg = await asyncio.wait_for(websocket.receive_json(), timeout=5.0) + except asyncio.TimeoutError: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + + if auth_msg.get("type") != "auth": + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + + token = auth_msg.get("access_token", "") + payload = decode_token(token) + if not payload or not getattr(payload, "sub", None): + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + + async with AsyncSessionLocal() as session: + result = await session.execute(select(User).where(User.id == payload.sub)) + user: User | None = result.scalar_one_or_none() + if not user or user.role not in (UserRole.admin, UserRole.superadmin): + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + + metrics_collector = websocket.app.state.metrics_collector + ws_manager = websocket.app.state.ws_manager + + try: + while True: + if metrics_collector: + metrics = metrics_collector.collect( + active_connections=ws_manager.active_connections_count if ws_manager else 0, + max_connections=getattr(ws_manager, "max_connections", 100), + ) + await websocket.send_json({"type": "metrics", "data": metrics}) + else: + await websocket.send_json( + {"type": "metrics", "data": {"detail": "Collector unavailable"}} + ) + await asyncio.sleep(5.0) + except WebSocketDisconnect: + pass + except Exception as exc: + logging.getLogger(__name__).error(f"Admin WS error: {exc}") + await websocket.close(code=status.WS_1011_INTERNAL_ERROR) diff --git a/api/v1/endpoints/asr_file.py b/api/v1/endpoints/asr_file.py index 43dea75..6eeb980 100644 --- a/api/v1/endpoints/asr_file.py +++ b/api/v1/endpoints/asr_file.py @@ -2,6 +2,12 @@ import logging from config import settings from fastapi import APIRouter, Depends, File, Form, UploadFile +from sqlalchemy.ext.asyncio import AsyncSession + +from core.deps import get_current_user_or_none +from db.session import get_db_session +from db.models import ASRSession +from db.enums import ASRSessionStatus, ASRSessionType from models.fast_api_models import PostFileRequest, V1BaseResponse, ASRData, RawData, SentencedData, DiarizedData from services.recognition_session import FileRecognitionSession from Recognizer.engine.file_recognition import process_file @@ -50,9 +56,24 @@ async def async_receive_file( params: PostFileRequest = Depends(get_file_request), recognizer: Recognizer = Depends(get_recognizer), punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), - diarizer: Diarizer = Depends(get_diarizer) + diarizer: Diarizer = Depends(get_diarizer), + current_user = Depends(get_current_user_or_none), + db: AsyncSession = Depends(get_db_session), ) -> V1BaseResponse: + from datetime import datetime, timezone + session = FileRecognitionSession(params=params) + asr_db_session = None + if current_user: + asr_db_session = ASRSession( + user_id=current_user.id, + session_type=ASRSessionType.file, + status=ASRSessionStatus.processing, + ) + db.add(asr_db_session) + await db.commit() + await db.refresh(asr_db_session) + try: await session.save_upload(file) logger.info(f"Получен и сохранён файл {file.filename}") @@ -63,6 +84,15 @@ async def async_receive_file( punctuator=punctuator, diarizer=diarizer ) + if asr_db_session: + try: + asr_db_session.status = ASRSessionStatus.completed if result_dict.get('success') else ASRSessionStatus.failed + asr_db_session.completed_at = datetime.now(timezone.utc) + asr_db_session.result_json = result_dict + await db.commit() + except Exception as exc: + logger.debug("Failed to save ASRSession result: %s", exc) + return V1BaseResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/api/v1/endpoints/asr_url.py b/api/v1/endpoints/asr_url.py index 0b83a6c..c2074df 100644 --- a/api/v1/endpoints/asr_url.py +++ b/api/v1/endpoints/asr_url.py @@ -6,6 +6,12 @@ from services.recognition_session import FileRecognitionSession from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from core.deps import get_current_user_or_none +from db.session import get_db_session +from db.models import ASRSession +from db.enums import ASRSessionStatus, ASRSessionType from Recognizer import get_recognizer, Recognizer from Recognizer.engine.file_recognition import process_file from Punctuation import get_punctuator, SbertPuncCaseOnnx @@ -22,10 +28,25 @@ async def post_v1( params: SyncASRRequest, recognizer: Recognizer = Depends(get_recognizer), punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), - diarizer: Diarizer = Depends(get_diarizer) + diarizer: Diarizer = Depends(get_diarizer), + current_user = Depends(get_current_user_or_none), + db: AsyncSession = Depends(get_db_session), ) -> V1BaseResponse: + from datetime import datetime, timezone + post_id = uuid.uuid4() session = FileRecognitionSession(post_id=str(post_id), params=params) + asr_db_session = None + if current_user: + asr_db_session = ASRSession( + user_id=current_user.id, + session_type=ASRSessionType.url, + status=ASRSessionStatus.processing, + ) + db.add(asr_db_session) + await db.commit() + await db.refresh(asr_db_session) + try: if params.AudioFileUrl: res, error_description, file_buffer = await getting_audiofile(params.AudioFileUrl, post_id) @@ -58,6 +79,15 @@ async def post_v1( punctuator=punctuator, diarizer=diarizer ) + if asr_db_session: + try: + asr_db_session.status = ASRSessionStatus.completed if result_dict.get('success') else ASRSessionStatus.failed + asr_db_session.completed_at = datetime.now(timezone.utc) + asr_db_session.result_json = result_dict + await db.commit() + except Exception as exc: + logger.debug("Failed to save ASRSession result: %s", exc) + return V1BaseResponse( success=result_dict.get('success', True), error_description=result_dict.get('error_description'), diff --git a/api/v1/endpoints/asr_ws.py b/api/v1/endpoints/asr_ws.py index 6bc487b..626128b 100644 --- a/api/v1/endpoints/asr_ws.py +++ b/api/v1/endpoints/asr_ws.py @@ -28,6 +28,11 @@ from services.asr_pipeline import process_audio_stream_chunk, process_final_audio from Recognizer import get_recognizer, Recognizer from Punctuation import get_punctuator, SbertPuncCaseOnnx +from db.session import get_db_session +from db.models import ASRSession +from db.enums import ASRSessionStatus, ASRSessionType +from sqlalchemy.ext.asyncio import AsyncSession +from datetime import datetime, timezone router = APIRouter(prefix="/asr", tags=["ASR"]) logger = logging.getLogger(__name__) @@ -52,6 +57,7 @@ async def websocket_endpoint( websocket: WebSocket, recognizer: Recognizer = Depends(get_recognizer), punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + db: AsyncSession = Depends(get_db_session), ): """ WebSocket endpoint для потокового распознавания речи (ASR). @@ -73,8 +79,48 @@ async def websocket_endpoint( logger.warning("Connection rejected for %s (max connections reached)", client_id) return + # 1a. Ожидание первого фрейма (auth или config) — таймаут 5 сек + user_id = None + pending_message = None + try: + auth_msg = await asyncio.wait_for(websocket.receive(), timeout=5.0) + if auth_msg.get("type") == "websocket.disconnect": + logger.info("Client %s disconnected before auth", client_id) + return + if auth_msg.get("text"): + import json + auth_data = json.loads(auth_msg["text"]) + if auth_data.get("type") == "auth": + token = auth_data.get("access_token", "") + from core.security import decode_token + payload = decode_token(token) + if payload and getattr(payload, "sub", None): + user_id = payload.sub + else: + # Первое сообщение не auth — сохраняем для обработки в цикле (config и т.д.) + pending_message = auth_msg + except asyncio.TimeoutError: + # Гостевой доступ: не закрываем соединение, просто логируем + logger.info("No auth for client %s, continuing as guest", client_id) + except Exception as exc: + logger.warning("Auth error for client %s: %s", client_id, exc) + # 2. Жизненный цикл сессии (гарантированная очистка в finally) async with audio_session_lifecycle(client_id) as session: + asr_db_session = None + if user_id: + session.user_id = user_id # привязка к пользователю (Этап 5) + # Создаём запись ASRSession в БД + asr_db_session = ASRSession( + user_id=user_id, + session_type=ASRSessionType.websocket, + status=ASRSessionStatus.processing, + request_ip=websocket.client.host if websocket.client else None, + ) + db.add(asr_db_session) + await db.commit() + await db.refresh(asr_db_session) + # 3. Регистрация хендлеров сообщений msg_router = MessageRouter() msg_router.register_handler(WSMessageType.config, handle_config) @@ -85,10 +131,14 @@ async def websocket_endpoint( while True: # 4. Получение сообщения с idle timeout try: - message = await asyncio.wait_for( - websocket.receive(), - timeout=settings.WS_IDLE_TIMEOUT_SEC, - ) + if pending_message is not None: + message = pending_message + pending_message = None + else: + message = await asyncio.wait_for( + websocket.receive(), + timeout=settings.WS_IDLE_TIMEOUT_SEC, + ) except asyncio.TimeoutError: logger.info("Idle timeout for client %s", client_id) await manager.send_message( @@ -108,6 +158,17 @@ async def websocket_endpoint( # Определяем тип содержимого: bytes (binary) или text (JSON) if message.get("bytes"): + if session.config is None: + logger.warning("Audio chunk received before config from %s", client_id) + await manager.send_message( + client_id, + WSErrorMessage( + code="missing_config", + message="Send config before audio chunks", + is_fatal=False, + ), + ) + continue # Binary frame: отправляем сырые байты напрямую в pipeline, без base64-обёртки await process_audio_stream_chunk( session, message["bytes"], recognizer, punctuator, manager, metrics @@ -133,13 +194,32 @@ async def websocket_endpoint( # Копирование флагов из конфига в сессию (для ASR pipeline) if isinstance(msg, WSConfigMessage): + session.config = msg session.wait_null_answers = msg.wait_null_answers session.do_dialogue = msg.do_dialogue session.do_punctuation = msg.do_punctuation session.channel_name = msg.channel_name or "Null" + logger.debug( + "Config set for %s: sample_rate=%d, format=%s, transport=%s", + client_id, + msg.sample_rate, + msg.audio_format, + msg.audio_transport, + ) # 6. Обработка аудио-чанка if isinstance(msg, WSAudioMessage): + if session.config is None: + logger.warning("Audio chunk received before config from %s", client_id) + await manager.send_message( + client_id, + WSErrorMessage( + code="missing_config", + message="Send config before audio chunks", + is_fatal=False, + ), + ) + continue chunk_bytes = b"" if msg.audio_base64: chunk_bytes = base64.b64decode(msg.audio_base64) @@ -157,6 +237,14 @@ async def websocket_endpoint( logger.info("Client %s disconnected normally", client_id) except Exception as exc: logger.exception("WS error for %s: %s", client_id, exc) + if asr_db_session: + try: + asr_db_session.status = ASRSessionStatus.failed + asr_db_session.error_message = str(exc) + asr_db_session.completed_at = datetime.now(timezone.utc) + await db.commit() + except Exception: + pass try: await manager.send_message( client_id, @@ -170,6 +258,15 @@ async def websocket_endpoint( pass finally: logger.info("Closing WS connection %s", client_id) + # Сохранение результата в БД + if asr_db_session: + try: + asr_db_session.status = ASRSessionStatus.completed + asr_db_session.completed_at = datetime.now(timezone.utc) + asr_db_session.result_json = session.ws_collected_asr_res + await db.commit() + except Exception as exc: + logger.debug("Failed to save ASRSession result: %s", exc) # Сохранение мета-информации в StateStore (аудит / восстановление) try: await state_store.set(f"session:{client_id}", session.to_dict()) diff --git a/api/v1/endpoints/auth.py b/api/v1/endpoints/auth.py new file mode 100644 index 0000000..9907942 --- /dev/null +++ b/api/v1/endpoints/auth.py @@ -0,0 +1,205 @@ +"""FastAPI-роутер аутентификации: JWT, Telegram, logout, change-password.""" + +from fastapi import APIRouter, Depends, HTTPException, Request, Response, status +from sqlalchemy.ext.asyncio import AsyncSession + +from config import settings + +from core.deps import get_current_user +from core.security import decode_token # type: ignore[import-untyped] +from db.models import User +from db.session import get_db_session +from models.auth import ( + ChangePasswordRequest, + TelegramAuthRequest, + TokenResponse, + UserLoginRequest, + UserRegisterRequest, +) +from services.auth_service import ( + authenticate_user, + blacklist_refresh_token, + generate_tokens, + get_or_create_telegram_user, + is_refresh_blacklisted, + register_user, + validate_telegram_init_data, +) + +router = APIRouter(prefix="/auth", tags=["auth"]) + + +@router.post("/register", response_model=TokenResponse) +async def auth_register( + payload: UserRegisterRequest, + response: Response, + db: AsyncSession = Depends(get_db_session), +): + """Регистрация нового пользователя. Refresh-токен в httpOnly cookie.""" + try: + user = await register_user(db, payload.email, payload.password, payload.full_name) + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) + + tokens = generate_tokens(user) + response.set_cookie( + key="refresh_token", + value=tokens["refresh_token"], + httponly=True, + secure=settings.IS_PROD, + samesite="strict", + max_age=7 * 24 * 3600, + ) + return TokenResponse(access_token=tokens["access_token"]) + + +@router.post("/login", response_model=TokenResponse) +async def auth_login( + payload: UserLoginRequest, + response: Response, + db: AsyncSession = Depends(get_db_session), +): + """Вход по email/пароль. Refresh-токен в httpOnly cookie.""" + user = await authenticate_user(db, payload.email, payload.password) + if not user: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Неверный email или пароль", + ) + + tokens = generate_tokens(user) + response.set_cookie( + key="refresh_token", + value=tokens["refresh_token"], + httponly=True, + secure=settings.IS_PROD, + samesite="strict", + max_age=7 * 24 * 3600, + ) + return TokenResponse(access_token=tokens["access_token"]) + + +@router.post("/refresh", response_model=TokenResponse) +async def auth_refresh( + request: Request, + response: Response, + db: AsyncSession = Depends(get_db_session), +): + """Обмен валидного refresh cookie на новую пару токенов.""" + refresh_token = request.cookies.get("refresh_token") + if not refresh_token: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Отсутствует refresh токен", + ) + + if await is_refresh_blacklisted(refresh_token): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Refresh токен отозван", + ) + + payload = decode_token(refresh_token) + if not payload or not getattr(payload, "sub", None): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Невалидный refresh токен", + ) + + from sqlalchemy import select + + result = await db.execute(select(User).where(User.id == payload.sub)) + user: User | None = result.scalar_one_or_none() + if not user or not user.is_active: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Пользователь не найден", + ) + + tokens = generate_tokens(user) + response.set_cookie( + key="refresh_token", + value=tokens["refresh_token"], + httponly=True, + secure=settings.IS_PROD, + samesite="strict", + max_age=7 * 24 * 3600, + ) + return TokenResponse(access_token=tokens["access_token"]) + + +@router.post("/logout") +async def auth_logout(request: Request, response: Response): + """Инвалидация refresh токена (blacklist) + очистка cookie.""" + refresh_token = request.cookies.get("refresh_token") + if refresh_token: + await blacklist_refresh_token(refresh_token, ttl_sec=7 * 24 * 3600) + response.delete_cookie(key="refresh_token", httponly=True, secure=settings.IS_PROD, samesite="strict") + return {"detail": "Успешный выход"} + + +@router.post("/change-password") +async def auth_change_password( + payload: ChangePasswordRequest, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Смена пароля текущего пользователя (требуется старый пароль).""" + from core.security import get_password_hash, verify_password # type: ignore[import-untyped] + + if not verify_password(payload.old_password, current_user.hashed_password): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Неверный текущий пароль", + ) + current_user.hashed_password = get_password_hash(payload.new_password) + await db.commit() + return {"detail": "Пароль успешно изменён"} + + +@router.get("/me") +async def auth_me(current_user: User = Depends(get_current_user)): + """Текущий аутентифицированный пользователь.""" + return { + "id": current_user.id, + "email": current_user.email, + "full_name": current_user.full_name, + "role": current_user.role, + "is_active": current_user.is_active, + } + + +@router.post("/telegram", response_model=TokenResponse) +async def auth_telegram( + payload: TelegramAuthRequest, + response: Response, + db: AsyncSession = Depends(get_db_session), +): + """Аутентификация через Telegram Web App (initData + HMAC-проверка).""" + import os + + bot_token = os.getenv("TELEGRAM_BOT_TOKEN", "") + if not bot_token: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="Telegram бот не настроен", + ) + + tg_data = validate_telegram_init_data(payload.init_data, bot_token) + if not tg_data: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Невалидная подпись Telegram initData", + ) + + user = await get_or_create_telegram_user(db, tg_data) + tokens = generate_tokens(user) + response.set_cookie( + key="refresh_token", + value=tokens["refresh_token"], + httponly=True, + secure=settings.IS_PROD, + samesite="strict", + max_age=7 * 24 * 3600, + ) + return TokenResponse(access_token=tokens["access_token"]) diff --git a/api/v1/endpoints/tg.py b/api/v1/endpoints/tg.py new file mode 100644 index 0000000..ff44912 --- /dev/null +++ b/api/v1/endpoints/tg.py @@ -0,0 +1,54 @@ +"""Точка входа для Telegram Web App.""" + +from fastapi import APIRouter, Request +from fastapi.responses import HTMLResponse + +router = APIRouter(tags=["telegram"]) + + +@router.get("/tg", response_class=HTMLResponse) +async def telegram_webapp_entry(request: Request): + """Возвращает HTML-страницу входа для Telegram Web App.""" + # В будущем заменить на TemplateResponse("tg/index.html", {"request": request}) + html = """ + + + + + ASR Telegram + + + + +

ASR Сервис

+

Авторизация...

+ + +""" + return HTMLResponse(content=html) diff --git a/api/v1/endpoints/user.py b/api/v1/endpoints/user.py new file mode 100644 index 0000000..f09d46b --- /dev/null +++ b/api/v1/endpoints/user.py @@ -0,0 +1,373 @@ +"""FastAPI-роутер пользовательского кабинета.""" + +import hashlib +import secrets +from datetime import datetime, timezone +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, status +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from core.deps import get_current_user +from db.enums import SubscriptionStatus +from db.models import ApiKey, ASRSession, Plan, Subscription, User +from db.session import get_db_session +from models.user import ( + ApiKeyCreateRequest, + ApiKeyCreateResponse, + ApiKeyResponse, + ASRSessionItem, + TelegramLinkStatus, + TelegramUnlinkRequest, + UserProfileResponse, + UserProfileUpdateRequest, + UserQuotaResponse, + UserStatsResponse, + UserSubscriptionResponse, +) + +router = APIRouter(prefix="/user", tags=["user"]) + + +@router.get("/profile", response_model=UserProfileResponse) +async def user_profile(current_user: User = Depends(get_current_user)): + """Получить профиль текущего пользователя.""" + return UserProfileResponse( + id=current_user.id, + email=current_user.email, + full_name=current_user.full_name, + phone=current_user.phone, + role=current_user.role, + is_active=current_user.is_active, + telegram_linked=current_user.telegram_id is not None, + ) + + +@router.put("/profile", response_model=UserProfileResponse) +async def user_profile_update( + payload: UserProfileUpdateRequest, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Обновить профиль текущего пользователя.""" + if payload.full_name is not None: + current_user.full_name = payload.full_name + if payload.phone is not None: + current_user.phone = payload.phone + await db.commit() + await db.refresh(current_user) + return UserProfileResponse( + id=current_user.id, + email=current_user.email, + full_name=current_user.full_name, + phone=current_user.phone, + role=current_user.role, + is_active=current_user.is_active, + telegram_linked=current_user.telegram_id is not None, + ) + + +@router.delete("/profile") +async def user_profile_delete( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Soft-delete аккаунта текущего пользователя.""" + current_user.is_active = False + await db.commit() + return {"detail": "Аккаунт деактивирован"} + + +@router.get("/quota", response_model=UserQuotaResponse) +async def user_quota( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Получить текущую квоту пользователя (заглушка до полной интеграции rate limiting).""" + result = await db.execute( + select(Subscription, Plan) + .join(Plan, Subscription.plan_id == Plan.id) + .where( + Subscription.user_id == current_user.id, + Subscription.status == SubscriptionStatus.active, + ) + ) + row = result.first() + if row: + sub, plan = row + return UserQuotaResponse( + plan_name=plan.name, + max_requests_per_minute=plan.max_requests_per_minute, + requests_used_this_minute=0, # TODO: интегрировать rate limiting счётчик + max_audio_duration_sec=plan.max_audio_duration_sec, + ) + return UserQuotaResponse() + + +@router.get("/subscription", response_model=UserSubscriptionResponse) +async def user_subscription( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Получить текущую подписку пользователя.""" + result = await db.execute( + select(Subscription, Plan) + .join(Plan, Subscription.plan_id == Plan.id) + .where(Subscription.user_id == current_user.id) + .order_by(Subscription.created_at.desc()) + ) + row = result.first() + if not row: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Подписка не найдена", + ) + sub, plan = row + return UserSubscriptionResponse( + status=sub.status, + plan_name=plan.name, + started_at=sub.started_at, + expires_at=sub.expires_at, + auto_renew=sub.auto_renew, + ) + + +@router.post("/subscription/upgrade") +async def user_subscription_upgrade( + current_user: User = Depends(get_current_user), +): + """Заглушка для запроса на смену тарифа.""" + return {"detail": "Функция смены тарифа в разработке"} + + +@router.post("/subscription/cancel") +async def user_subscription_cancel( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Отмена авто-продления текущей подписки.""" + result = await db.execute( + select(Subscription).where( + Subscription.user_id == current_user.id, + Subscription.status == SubscriptionStatus.active, + ) + ) + sub: Optional[Subscription] = result.scalar_one_or_none() + if not sub: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Активная подписка не найдена", + ) + sub.auto_renew = False + await db.commit() + return {"detail": "Авто-продление отменено"} + + +@router.get("/sessions", response_model=list[ASRSessionItem]) +async def user_sessions( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), + limit: int = 20, + offset: int = 0, +): + """Список ASR-сессий текущего пользователя.""" + result = await db.execute( + select(ASRSession) + .where(ASRSession.user_id == current_user.id) + .order_by(ASRSession.created_at.desc()) + .limit(limit) + .offset(offset) + ) + sessions = result.scalars().all() + return [ + ASRSessionItem( + id=s.id, + session_type=s.session_type, + status=s.status, + audio_duration_sec=s.audio_duration_sec, + created_at=s.created_at, + completed_at=s.completed_at, + ) + for s in sessions + ] + + +@router.get("/sessions/{session_id}") +async def user_session_detail( + session_id: str, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Детали конкретной ASR-сессии.""" + result = await db.execute( + select(ASRSession).where( + ASRSession.id == session_id, + ASRSession.user_id == current_user.id, + ) + ) + session: Optional[ASRSession] = result.scalar_one_or_none() + if not session: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Сессия не найдена", + ) + return { + "id": session.id, + "session_type": session.session_type, + "status": session.status, + "audio_duration_sec": session.audio_duration_sec, + "processing_duration_sec": session.processing_duration_sec, + "cost": float(session.cost) if session.cost is not None else None, + "result_json": session.result_json, + "created_at": session.created_at, + "completed_at": session.completed_at, + "error_message": session.error_message, + } + + +@router.get("/stats", response_model=UserStatsResponse) +async def user_stats( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Агрегированная статистика пользователя.""" + total_result = await db.execute( + select(func.count(ASRSession.id)).where(ASRSession.user_id == current_user.id) + ) + total_sessions = total_result.scalar() or 0 + + audio_result = await db.execute( + select(func.coalesce(func.sum(ASRSession.audio_duration_sec), 0)).where( + ASRSession.user_id == current_user.id + ) + ) + total_audio_sec = audio_result.scalar() or 0 + + now = datetime.now(timezone.utc) + month_start = now.replace(day=1, hour=0, minute=0, second=0, microsecond=0) + month_result = await db.execute( + select(func.count(ASRSession.id)).where( + ASRSession.user_id == current_user.id, + ASRSession.created_at >= month_start, + ) + ) + sessions_this_month = month_result.scalar() or 0 + + return UserStatsResponse( + total_sessions=total_sessions, + total_audio_hours=round(total_audio_sec / 3600, 2), + sessions_this_month=sessions_this_month, + ) + + +@router.get("/api-keys", response_model=list[ApiKeyResponse]) +async def user_api_keys( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Список API-ключей пользователя (без plain key).""" + result = await db.execute( + select(ApiKey) + .where(ApiKey.user_id == current_user.id) + .order_by(ApiKey.created_at.desc()) + ) + keys = result.scalars().all() + return [ + ApiKeyResponse( + id=k.id, + name=k.name, + is_active=k.is_active, + created_at=k.created_at, + last_used_at=k.last_used_at, + ) + for k in keys + ] + + +@router.post("/api-keys", response_model=ApiKeyCreateResponse) +async def user_api_key_create( + payload: ApiKeyCreateRequest, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Создать новый API-ключ. Plain key возвращается только один раз.""" + plain_key = f"asr_{secrets.token_urlsafe(32)}" + key_hash = hashlib.sha256(plain_key.encode()).hexdigest() + + api_key = ApiKey( + user_id=current_user.id, + name=payload.name, + key_hash=key_hash, + is_active=True, + ) + db.add(api_key) + await db.commit() + await db.refresh(api_key) + + return ApiKeyCreateResponse( + id=api_key.id, + name=api_key.name, + is_active=api_key.is_active, + created_at=api_key.created_at, + last_used_at=api_key.last_used_at, + plain_key=plain_key, + ) + + +@router.delete("/api-keys/{key_id}") +async def user_api_key_delete( + key_id: str, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Отозвать API-ключ.""" + result = await db.execute( + select(ApiKey).where( + ApiKey.id == key_id, + ApiKey.user_id == current_user.id, + ) + ) + key: Optional[ApiKey] = result.scalar_one_or_none() + if not key: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Ключ не найден", + ) + key.is_active = False + await db.commit() + return {"detail": "Ключ отозван"} + + +@router.get("/telegram/link", response_model=TelegramLinkStatus) +async def user_telegram_link( + current_user: User = Depends(get_current_user), +): + """Проверить, привязан ли Telegram-аккаунт.""" + return TelegramLinkStatus( + linked=current_user.telegram_id is not None, + telegram_username=current_user.telegram_username, + ) + + +@router.post("/telegram/unlink") +async def user_telegram_unlink( + payload: TelegramUnlinkRequest, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db_session), +): + """Отвязать Telegram-аккаунт от пользователя.""" + if not current_user.hashed_password: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Нельзя отвязать Telegram для аккаунта без пароля. Установите пароль.", + ) + + current_user.telegram_id = None + current_user.telegram_username = None + current_user.telegram_first_name = None + current_user.telegram_last_name = None + current_user.telegram_photo_url = None + current_user.telegram_auth_date = None + await db.commit() + return {"detail": "Telegram отвязан"} diff --git a/core/deps.py b/core/deps.py new file mode 100644 index 0000000..497e57c --- /dev/null +++ b/core/deps.py @@ -0,0 +1,144 @@ +"""FastAPI-зависимости для аутентификации и авторизации (RBAC + API Key).""" + +from fastapi import Depends, HTTPException, Request, status +from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from db.enums import UserRole +from db.models import ApiKey, User +from db.session import get_db_session + +# Предполагается, что в core/security.py есть decode_token и TokenPayload +from core.security import decode_token # type: ignore[import-untyped] + +bearer_scheme = HTTPBearer(auto_error=False) + + +async def _extract_token( + request: Request, + credentials: HTTPAuthorizationCredentials | None, +) -> str | None: + """Извлекает токен из заголовка Authorization, cookie или query-параметра.""" + if credentials: + return credentials.credentials + token = request.cookies.get("access_token") + if token: + return token + return request.query_params.get("token") + + +async def get_current_user( + request: Request, + credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme), + db: AsyncSession = Depends(get_db_session), +) -> User: + """Возвращает текущего пользователя по JWT access token.""" + token = await _extract_token(request, credentials) + if not token: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Не предоставлен токен авторизации", + ) + + payload = decode_token(token) + if not payload or not getattr(payload, "sub", None): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Невалидный или просроченный токен", + ) + + result = await db.execute(select(User).where(User.id == payload.sub)) + user: User | None = result.scalar_one_or_none() + if not user or not user.is_active: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Пользователь не найден или деактивирован", + ) + return user + + +def require_admin(user: User = Depends(get_current_user)) -> User: + """Проверяет, что пользователь — админ или суперадмин.""" + if user.role not in (UserRole.admin, UserRole.superadmin): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Требуются права администратора", + ) + return user + + +def require_superadmin(user: User = Depends(get_current_user)) -> User: + """Проверяет, что пользователь — суперадмин.""" + if user.role != UserRole.superadmin: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Требуются права суперадминистратора", + ) + return user + + +async def get_current_user_or_none( + request: Request, + credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme), + db: AsyncSession = Depends(get_db_session), +) -> User | None: + """Возвращает пользователя или None (для опциональной авторизации).""" + try: + return await get_current_user(request, credentials, db) + except HTTPException: + return None + + +async def require_api_key( + request: Request, + db: AsyncSession = Depends(get_db_session), +) -> User: + """Аутентификация по API-ключу из заголовка X-API-Key.""" + api_key = request.headers.get("X-API-Key") + if not api_key: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Не предоставлен API-ключ", + ) + + import hashlib + from datetime import datetime, timezone + + key_hash = hashlib.sha256(api_key.encode()).hexdigest() + + result = await db.execute( + select(ApiKey).where(ApiKey.key_hash == key_hash, ApiKey.is_active.is_(True)) + ) + key_obj: ApiKey | None = result.scalar_one_or_none() + if not key_obj: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Невалидный или отозванный API-ключ", + ) + + # Обновляем last_used_at + key_obj.last_used_at = datetime.now(timezone.utc) + await db.commit() + + result_user = await db.execute(select(User).where(User.id == key_obj.user_id)) + user: User | None = result_user.scalar_one_or_none() + if not user or not user.is_active: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Пользователь ключа не найден или деактивирован", + ) + return user + + +async def require_api_key_or_auth( + request: Request, + credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme), + db: AsyncSession = Depends(get_db_session), +) -> User: + """Пробует JWT-аутентификацию, затем API-ключ.""" + try: + return await get_current_user(request, credentials, db) + except HTTPException: + pass + return await require_api_key(request, db) diff --git a/core/middleware.py b/core/middleware.py new file mode 100644 index 0000000..5727ab5 --- /dev/null +++ b/core/middleware.py @@ -0,0 +1,84 @@ +"""Дополнительные middleware: rate limiting, maintenance mode.""" + +import time +from typing import Optional + +from fastapi import Request, Response +from starlette.middleware.base import BaseHTTPMiddleware + +from services.admin_service import is_maintenance_mode + + +class RateLimitMiddleware(BaseHTTPMiddleware): + """Rate limiting по количеству запросов в минуту (in-memory).""" + + def __init__( + self, + app, + max_requests: int = 60, + window_seconds: float = 60.0, + exempt_paths: Optional[set[str]] = None, + ): + super().__init__(app) + self.max_requests = max_requests + self.window_seconds = window_seconds + self.exempt_paths = exempt_paths or { + "/docs", + "/openapi.json", + "/static", + "/admin/login", + "/api/v1/auth/login", + "/api/v1/auth/register", + "/api/v1/auth/refresh", + "/api/v1/auth/telegram", + "/tg", + } + self._storage: dict[str, list[float]] = {} + + async def dispatch(self, request: Request, call_next): + path = request.url.path + + # Maintenance mode для изменяющих запросов + if is_maintenance_mode() and request.method not in ("GET", "HEAD", "OPTIONS"): + return Response( + content='{"detail":"Сервис на обслуживании"}', + status_code=503, + media_type="application/json", + ) + + # Пропускаем exempt пути + if any(path.startswith(ep) for ep in self.exempt_paths): + return await call_next(request) + + # Определяем ключ лимита + api_key = request.headers.get("X-API-Key") + auth_header = request.headers.get("Authorization", "") + client_ip = request.client.host if request.client else "unknown" + + if api_key: + import hashlib + + key = f"rate_limit:apikey:{hashlib.sha256(api_key.encode()).hexdigest()[:16]}" + elif auth_header.startswith("Bearer "): + token = auth_header[7:] + key = f"rate_limit:jwt:{token[:16]}" + else: + key = f"rate_limit:ip:{client_ip}" + + now = time.time() + timestamps = self._storage.get(key, []) + timestamps = [t for t in timestamps if now - t < self.window_seconds] + + if len(timestamps) >= self.max_requests: + retry_after = int(self.window_seconds - (now - timestamps[0])) + return Response( + content=f'{{"detail":"Превышен лимит запросов","retry_after":{retry_after}}}', + status_code=429, + media_type="application/json", + headers={"Retry-After": str(retry_after)}, + ) + + timestamps.append(now) + self._storage[key] = timestamps + + return await call_next(request) diff --git a/db/__init__.py b/db/__init__.py new file mode 100644 index 0000000..2a80373 --- /dev/null +++ b/db/__init__.py @@ -0,0 +1 @@ +"""Пакет инфраструктуры базы данных (SQLAlchemy 2.0 + Alembic).""" diff --git a/db/base.py b/db/base.py new file mode 100644 index 0000000..d45ff76 --- /dev/null +++ b/db/base.py @@ -0,0 +1,24 @@ +"""Базовый декларативный класс SQLAlchemy 2.0.""" + +import uuid +from datetime import datetime, timezone + +from sqlalchemy import String +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column + + +class Base(DeclarativeBase): + """Базовый класс с UUID-PK и timestamp-полями для всех моделей.""" + + id: Mapped[str] = mapped_column( + String(36), + primary_key=True, + default=lambda: str(uuid.uuid4()), + ) + created_at: Mapped[datetime] = mapped_column( + default=lambda: datetime.now(timezone.utc), + ) + updated_at: Mapped[datetime] = mapped_column( + default=lambda: datetime.now(timezone.utc), + onupdate=lambda: datetime.now(timezone.utc), + ) diff --git a/db/enums.py b/db/enums.py new file mode 100644 index 0000000..a2f1676 --- /dev/null +++ b/db/enums.py @@ -0,0 +1,51 @@ +"""Перечисления (Enum) для SQLAlchemy-моделей.""" + +import enum + + +class UserRole(str, enum.Enum): + """Роли пользователей.""" + + user = "user" + admin = "admin" + superadmin = "superadmin" + + +class SubscriptionStatus(str, enum.Enum): + """Статусы подписки.""" + + active = "active" + expired = "expired" + cancelled = "cancelled" + + +class TransactionStatus(str, enum.Enum): + """Статусы платежной транзакции.""" + + pending = "pending" + succeeded = "succeeded" + cancelled = "cancelled" + + +class ASRSessionStatus(str, enum.Enum): + """Статусы сессии распознавания.""" + + processing = "processing" + completed = "completed" + failed = "failed" + + +class ASRSessionType(str, enum.Enum): + """Типы сессии распознавания.""" + + url = "url" + file = "file" + websocket = "websocket" + + +class SystemLogLevel(str, enum.Enum): + """Уровни системного лога.""" + + info = "info" + warning = "warning" + error = "error" diff --git a/db/models.py b/db/models.py new file mode 100644 index 0000000..c0eeb9a --- /dev/null +++ b/db/models.py @@ -0,0 +1,228 @@ +"""SQLAlchemy 2.0 ORM-модели для ASR-сервиса.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from typing import Optional + +from sqlalchemy import BigInteger, JSON, ForeignKey, Integer, Numeric, String, Text +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from db.base import Base +from db.enums import ( + ASRSessionStatus, + ASRSessionType, + SubscriptionStatus, + SystemLogLevel, + TransactionStatus, + UserRole, +) + + +class User(Base): + """Пользователь системы.""" + + __tablename__ = "users" + + email: Mapped[str] = mapped_column(String(255), unique=True, index=True) + hashed_password: Mapped[str] = mapped_column(Text) + full_name: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + phone: Mapped[Optional[str]] = mapped_column(String(50), nullable=True) + role: Mapped[UserRole] = mapped_column(String(20), default=UserRole.user) + is_active: Mapped[bool] = mapped_column(default=True) + email_verified: Mapped[bool] = mapped_column(default=False) + last_login_at: Mapped[Optional[datetime]] = mapped_column(nullable=True) + + # Telegram Web App fields + telegram_id: Mapped[Optional[int]] = mapped_column( + BigInteger, unique=True, nullable=True, index=True + ) + telegram_username: Mapped[Optional[str]] = mapped_column( + String(255), nullable=True + ) + telegram_first_name: Mapped[Optional[str]] = mapped_column( + String(255), nullable=True + ) + telegram_last_name: Mapped[Optional[str]] = mapped_column( + String(255), nullable=True + ) + telegram_photo_url: Mapped[Optional[str]] = mapped_column( + Text, nullable=True + ) + telegram_auth_date: Mapped[Optional[datetime]] = mapped_column( + nullable=True + ) + + # relationships + subscriptions: Mapped[list["Subscription"]] = relationship( + back_populates="user", lazy="selectin" + ) + api_keys: Mapped[list["ApiKey"]] = relationship( + back_populates="user", lazy="selectin" + ) + asr_sessions: Mapped[list["ASRSession"]] = relationship( + back_populates="user", lazy="selectin" + ) + admin_audit_logs: Mapped[list["AdminAuditLog"]] = relationship( + back_populates="admin", lazy="selectin" + ) + + +class Plan(Base): + """Тарифный план (справочник).""" + + __tablename__ = "plans" + + code: Mapped[str] = mapped_column(String(50), unique=True) + name: Mapped[str] = mapped_column(String(255)) + description: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + max_requests_per_minute: Mapped[int] = mapped_column(Integer, default=60) + max_audio_duration_sec: Mapped[Optional[int]] = mapped_column( + Integer, nullable=True + ) + price_per_month: Mapped[Optional[int]] = mapped_column( + Numeric(10, 2), nullable=True + ) + is_active: Mapped[bool] = mapped_column(default=True) + + subscriptions: Mapped[list["Subscription"]] = relationship( + back_populates="plan", lazy="selectin" + ) + + +class Subscription(Base): + """Подписка пользователя.""" + + __tablename__ = "subscriptions" + + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id"), index=True + ) + plan_id: Mapped[str] = mapped_column( + String(36), ForeignKey("plans.id"), index=True + ) + status: Mapped[SubscriptionStatus] = mapped_column( + String(20), default=SubscriptionStatus.active + ) + started_at: Mapped[datetime] = mapped_column( + default=lambda: datetime.now(timezone.utc) + ) + expires_at: Mapped[Optional[datetime]] = mapped_column(nullable=True) + auto_renew: Mapped[bool] = mapped_column(default=False) + yookassa_payment_id: Mapped[Optional[str]] = mapped_column( + String(255), nullable=True + ) + + user: Mapped["User"] = relationship(back_populates="subscriptions") + plan: Mapped["Plan"] = relationship(back_populates="subscriptions") + transactions: Mapped[list["Transaction"]] = relationship( + back_populates="subscription", lazy="selectin" + ) + + +class Transaction(Base): + """Платёжная транзакция.""" + + __tablename__ = "transactions" + + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id"), index=True + ) + subscription_id: Mapped[Optional[str]] = mapped_column( + String(36), ForeignKey("subscriptions.id"), nullable=True, index=True + ) + amount: Mapped[Optional[int]] = mapped_column(Numeric(10, 2), nullable=True) + currency: Mapped[str] = mapped_column(String(3), default="RUB") + status: Mapped[TransactionStatus] = mapped_column( + String(20), default=TransactionStatus.pending + ) + payment_provider: Mapped[str] = mapped_column(String(50), default="yookassa") + external_payment_id: Mapped[Optional[str]] = mapped_column( + String(255), nullable=True, index=True + ) + meta: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True) + + subscription: Mapped[Optional["Subscription"]] = relationship( + back_populates="transactions" + ) + + +class ASRSession(Base): + """Сессия распознавания речи.""" + + __tablename__ = "asr_sessions" + + user_id: Mapped[Optional[str]] = mapped_column( + String(36), ForeignKey("users.id"), nullable=True, index=True + ) + session_type: Mapped[ASRSessionType] = mapped_column(String(20)) + status: Mapped[ASRSessionStatus] = mapped_column( + String(20), default=ASRSessionStatus.processing + ) + audio_duration_sec: Mapped[Optional[float]] = mapped_column(nullable=True) + processing_duration_sec: Mapped[Optional[float]] = mapped_column(nullable=True) + cost: Mapped[Optional[int]] = mapped_column(Numeric(10, 2), nullable=True) + result_json: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True) + completed_at: Mapped[Optional[datetime]] = mapped_column(nullable=True) + error_message: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + request_ip: Mapped[Optional[str]] = mapped_column(String(45), nullable=True) + user_agent: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + + user: Mapped[Optional["User"]] = relationship(back_populates="asr_sessions") + + +class ApiKey(Base): + """API-ключ для программного доступа.""" + + __tablename__ = "api_keys" + + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id"), index=True + ) + name: Mapped[str] = mapped_column(String(255)) + key_hash: Mapped[str] = mapped_column(String(255), unique=True, index=True) + permissions: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True) + is_active: Mapped[bool] = mapped_column(default=True) + rate_limit_override: Mapped[Optional[int]] = mapped_column( + Integer, nullable=True + ) + last_used_at: Mapped[Optional[datetime]] = mapped_column(nullable=True) + + user: Mapped["User"] = relationship(back_populates="api_keys") + + +class SystemLog(Base): + """Системный лог.""" + + __tablename__ = "system_logs" + + level: Mapped[SystemLogLevel] = mapped_column(String(20), index=True) + component: Mapped[str] = mapped_column(String(100), index=True) + message: Mapped[str] = mapped_column(Text) + meta: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True) + + +class AdminAuditLog(Base): + """Аудит действий администраторов.""" + + __tablename__ = "admin_audit_logs" + + admin_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id"), index=True + ) + action: Mapped[str] = mapped_column(String(100)) + target_type: Mapped[Optional[str]] = mapped_column(String(50), nullable=True) + target_id: Mapped[Optional[str]] = mapped_column(String(36), nullable=True) + details: Mapped[Optional[dict]] = mapped_column(JSON, nullable=True) + + admin: Mapped["User"] = relationship(back_populates="admin_audit_logs") + + +class TelegramBotConfig(Base): + """Конфигурация Telegram-бота (опционально, для админ-панели).""" + + __tablename__ = "telegram_bot_configs" + + bot_token_hash: Mapped[str] = mapped_column(Text) + webapp_url: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + is_active: Mapped[bool] = mapped_column(default=True) diff --git a/db/session.py b/db/session.py new file mode 100644 index 0000000..0d72896 --- /dev/null +++ b/db/session.py @@ -0,0 +1,40 @@ +"""Асинхронная сессия SQLAlchemy для FastAPI.""" + +import os + +from sqlalchemy import NullPool +from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine + +# Placeholder: при запуске приложения URL должен переопределяться +# через config.py (Settings.DATABASE_URL). +# Для локальной разработки используем async SQLite (aiosqlite). +DATABASE_URL = os.getenv( + "DATABASE_URL", + "sqlite+aiosqlite:///./asr_local.db", +) + +# Для SQLite рекомендуется NullPool, чтобы избежать проблем с пулом +# в однопоточном режиме. +_pool_cls = NullPool if DATABASE_URL.startswith("sqlite") else None + +engine = create_async_engine( + DATABASE_URL, + echo=False, + future=True, + poolclass=_pool_cls, +) +AsyncSessionLocal = async_sessionmaker( + bind=engine, + expire_on_commit=False, + autoflush=False, + autocommit=False, +) + + +async def get_db_session(): + """Генератор сессии для FastAPI Depends.""" + async with AsyncSessionLocal() as session: + try: + yield session + finally: + await session.close() diff --git a/main.py b/main.py index b10667a..d986908 100644 --- a/main.py +++ b/main.py @@ -1,3 +1,4 @@ +import asyncio import logging import time import uvicorn @@ -23,6 +24,12 @@ from routes.ws_audio_transkrib import router as ws_audio_transkrib_router from api.legacy import router as legacy_router from api.v1.api import router as api_v1_router +from api.v1.endpoints.tg import router as tg_router +from routes.admin import router as admin_html_router +from routes.user import router as user_html_router +from core.middleware import RateLimitMiddleware +from api.v1.endpoints.admin_ws import router as admin_ws_router +from services.metrics_reporter import metrics_reporter_loop import models from config import WS_DESCRIPTION @@ -105,6 +112,11 @@ async def lifespan(app): interval_sec=settings.WS_STATUS_BROADCAST_INTERVAL_SEC, ) + # Запуск фоновой задачи записи метрик в БД (Этап 5) + app.state.metrics_task = asyncio.create_task( + metrics_reporter_loop(app.state, interval_sec=30.0) + ) + # Настройка сборщика мусора. gc.set_threshold(500, 5, 5) @@ -144,6 +156,14 @@ async def lifespan(app): yield # Здесь приложение работает + # Остановка фоновой задачи метрик + if hasattr(app.state, "metrics_task"): + app.state.metrics_task.cancel() + try: + await app.state.metrics_task + except asyncio.CancelledError: + pass + # Graceful shutdown WebSocket (Задача 6.4, 6.9) if hasattr(app.state, "ws_manager"): app.state.ws_manager.stop_status_broadcast() @@ -205,6 +225,13 @@ async def lifespan(app): minimum_size=500 ) +# Rate limiting middleware (Этап 5) +app.add_middleware( + RateLimitMiddleware, + max_requests=60, + window_seconds=60.0, +) + # Exception handlers register_exception_handlers(app) @@ -215,6 +242,10 @@ async def lifespan(app): app.include_router(ws_audio_transkrib_router, tags=["legacy"]) app.include_router(legacy_router, tags=["legacy"]) app.include_router(api_v1_router, tags=["api/v1"]) +app.include_router(tg_router) +app.include_router(admin_html_router) +app.include_router(admin_ws_router) +app.include_router(user_html_router) try: if __name__ == '__main__': diff --git a/models/admin.py b/models/admin.py new file mode 100644 index 0000000..a841ea6 --- /dev/null +++ b/models/admin.py @@ -0,0 +1,162 @@ +"""Pydantic-модели для админ-панели.""" + +from datetime import datetime +from typing import Any, Optional + +from pydantic import BaseModel, Field + + +class PaginationParams(BaseModel): + """Параметры пагинации для списков.""" + + page: int = Field(1, ge=1) + per_page: int = Field(20, ge=1, le=100) + + +class AdminUserListItem(BaseModel): + """Элемент списка пользователей в админке.""" + + id: str + email: str + full_name: Optional[str] = None + role: str + is_active: bool + created_at: Optional[datetime] = None + last_login_at: Optional[datetime] = None + telegram_linked: bool = False + + +class AdminUserDetailResponse(AdminUserListItem): + """Детальная информация о пользователе.""" + + phone: Optional[str] = None + telegram_id: Optional[int] = None + telegram_username: Optional[str] = None + + +class AdminUserUpdateRequest(BaseModel): + """Запрос на обновление пользователя админом.""" + + full_name: Optional[str] = None + phone: Optional[str] = None + role: Optional[str] = None + is_active: Optional[bool] = None + + +class AdminPlanCreateUpdateRequest(BaseModel): + """Запрос на создание/обновление тарифа.""" + + code: str + name: str + description: Optional[str] = None + max_requests_per_minute: int = 60 + max_audio_duration_sec: Optional[int] = None + price_per_month: Optional[float] = None + is_active: bool = True + + +class AdminPlanResponse(AdminPlanCreateUpdateRequest): + """Ответ с тарифом.""" + + id: str + created_at: Optional[datetime] = None + updated_at: Optional[datetime] = None + + +class AdminSubscriptionResponse(BaseModel): + """Ответ с подпиской.""" + + id: str + user_id: str + plan_id: str + plan_name: Optional[str] = None + status: str + started_at: Optional[datetime] = None + expires_at: Optional[datetime] = None + auto_renew: bool = False + + +class AdminTransactionResponse(BaseModel): + """Ответ с транзакцией.""" + + id: str + user_id: str + subscription_id: Optional[str] = None + amount: Optional[float] = None + currency: str + status: str + payment_provider: str + external_payment_id: Optional[str] = None + created_at: Optional[datetime] = None + + +class AdminSystemLogResponse(BaseModel): + """Ответ с системным логом.""" + + id: str + level: str + component: str + message: str + meta: Optional[dict] = None + created_at: Optional[datetime] = None + + +class AdminAuditLogResponse(BaseModel): + """Ответ с аудит-логом.""" + + id: str + admin_id: str + action: str + target_type: Optional[str] = None + target_id: Optional[str] = None + details: Optional[dict] = None + created_at: Optional[datetime] = None + + +class AdminApiKeyResponse(BaseModel): + """Ответ с API-ключом (админка).""" + + id: str + user_id: str + user_email: Optional[str] = None + name: str + is_active: bool + created_at: Optional[datetime] = None + last_used_at: Optional[datetime] = None + + +class AdminMetricsResponse(BaseModel): + """Ответ с текущими метриками системы.""" + + active_connections: int = 0 + active_tasks: int = 0 + cpu_percent: Optional[float] = None + gpu_utilization: Optional[float] = None + queue_depth: int = 0 + uptime_seconds: float = 0.0 + + +class AdminMaintenanceToggle(BaseModel): + """Запрос на переключение режима обслуживания.""" + + enabled: bool + + +class AdminTelegramConfig(BaseModel): + """Конфигурация Telegram-бота.""" + + webapp_url: Optional[str] = None + is_active: bool = True + + +class AdminTelegramStats(BaseModel): + """Статистика Telegram Web App.""" + + total_telegram_users: int = 0 + total_webapp_sessions: int = 0 + + +class AdminBroadcastRequest(BaseModel): + """Запрос на рассылку сообщения.""" + + message: str = Field(..., min_length=1, max_length=4096) diff --git a/models/auth.py b/models/auth.py new file mode 100644 index 0000000..f03d64a --- /dev/null +++ b/models/auth.py @@ -0,0 +1,38 @@ +"""Pydantic-модели для аутентификации и авторизации.""" + +from pydantic import BaseModel, EmailStr, Field + + +class UserRegisterRequest(BaseModel): + """Запрос на регистрацию по email/пароль.""" + + email: EmailStr + password: str = Field(..., min_length=6) + full_name: str | None = None + + +class UserLoginRequest(BaseModel): + """Запрос на вход по email/пароль.""" + + email: EmailStr + password: str + + +class TokenResponse(BaseModel): + """Ответ с access-токеном (refresh — в httpOnly cookie).""" + + access_token: str + token_type: str = "bearer" + + +class ChangePasswordRequest(BaseModel): + """Запрос на смену пароля.""" + + old_password: str + new_password: str = Field(..., min_length=6) + + +class TelegramAuthRequest(BaseModel): + """Запрос на аутентификацию через Telegram Web App.""" + + init_data: str diff --git a/models/user.py b/models/user.py new file mode 100644 index 0000000..1ccb3b0 --- /dev/null +++ b/models/user.py @@ -0,0 +1,98 @@ +"""Pydantic-модели для пользовательского кабинета.""" + +from datetime import datetime +from typing import Optional + +from pydantic import BaseModel + + +class UserProfileResponse(BaseModel): + """Ответ с профилем пользователя.""" + + id: str + email: str + full_name: Optional[str] = None + phone: Optional[str] = None + role: str + is_active: bool + telegram_linked: bool = False + + +class UserProfileUpdateRequest(BaseModel): + """Запрос на обновление профиля.""" + + full_name: Optional[str] = None + phone: Optional[str] = None + + +class UserQuotaResponse(BaseModel): + """Ответ с квотой пользователя.""" + + plan_name: Optional[str] = None + max_requests_per_minute: int = 0 + requests_used_this_minute: int = 0 + max_audio_duration_sec: Optional[int] = None + + +class UserSubscriptionResponse(BaseModel): + """Ответ с данными подписки.""" + + status: str + plan_name: Optional[str] = None + started_at: Optional[datetime] = None + expires_at: Optional[datetime] = None + auto_renew: bool = False + + +class ASRSessionItem(BaseModel): + """Элемент списка ASR-сессий.""" + + id: str + session_type: str + status: str + audio_duration_sec: Optional[float] = None + created_at: Optional[datetime] = None + completed_at: Optional[datetime] = None + + +class UserStatsResponse(BaseModel): + """Агрегированная статистика пользователя.""" + + total_sessions: int = 0 + total_audio_hours: float = 0.0 + sessions_this_month: int = 0 + + +class ApiKeyResponse(BaseModel): + """Ответ с данными API-ключа (без plain key).""" + + id: str + name: str + is_active: bool + created_at: Optional[datetime] = None + last_used_at: Optional[datetime] = None + + +class ApiKeyCreateRequest(BaseModel): + """Запрос на создание API-ключа.""" + + name: str + + +class ApiKeyCreateResponse(ApiKeyResponse): + """Ответ при создании ключа (с plain key, один раз).""" + + plain_key: str + + +class TelegramLinkStatus(BaseModel): + """Статус привязки Telegram.""" + + linked: bool + telegram_username: Optional[str] = None + + +class TelegramUnlinkRequest(BaseModel): + """Запрос на отвязку Telegram.""" + + password: Optional[str] = None diff --git a/routes/admin.py b/routes/admin.py new file mode 100644 index 0000000..d771a8e --- /dev/null +++ b/routes/admin.py @@ -0,0 +1,108 @@ +"""HTML-роуты админ-панели (Jinja2). + +TODO: подключить router в main.py через app.include_router(routes.admin.router). +""" + +from fastapi import APIRouter, Depends, Request +from fastapi.templating import Jinja2Templates + +from core.deps import require_admin + +templates = Jinja2Templates(directory="templates") +router = APIRouter(prefix="/admin") + + +@router.get("/dashboard") +async def admin_dashboard_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница dashboard админ-панели.""" + return templates.TemplateResponse("admin/dashboard.html", {"request": request}) + + +@router.get("/login") +async def admin_login_page(request: Request): + """Страница входа в админ-панель.""" + return templates.TemplateResponse("admin/login.html", {"request": request}) + + +@router.get("/users") +async def admin_users_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница управления пользователями.""" + return templates.TemplateResponse("admin/users.html", {"request": request}) + + +@router.get("/sessions") +async def admin_sessions_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница мониторинга сессий.""" + return templates.TemplateResponse("admin/sessions.html", {"request": request}) + + +@router.get("/subscriptions") +async def admin_subscriptions_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница управления подписками.""" + return templates.TemplateResponse("admin/subscriptions.html", {"request": request}) + + +@router.get("/transactions") +async def admin_transactions_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница транзакций.""" + return templates.TemplateResponse("admin/transactions.html", {"request": request}) + + +@router.get("/tariffs") +async def admin_tariffs_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница тарифных планов.""" + return templates.TemplateResponse("admin/tariffs.html", {"request": request}) + + +@router.get("/api-keys") +async def admin_api_keys_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница управления API-ключами.""" + return templates.TemplateResponse("admin/api_keys.html", {"request": request}) + + +@router.get("/logs") +async def admin_logs_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница системных логов.""" + return templates.TemplateResponse("admin/logs.html", {"request": request}) + + +@router.get("/settings") +async def admin_settings_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница настроек и maintenance mode.""" + return templates.TemplateResponse("admin/settings.html", {"request": request}) + + +@router.get("/telegram") +async def admin_telegram_page( + request: Request, + _admin=Depends(require_admin), +): + """Страница настроек Telegram-интеграции.""" + return templates.TemplateResponse("admin/telegram.html", {"request": request}) diff --git a/routes/user.py b/routes/user.py new file mode 100644 index 0000000..fe3df56 --- /dev/null +++ b/routes/user.py @@ -0,0 +1,84 @@ +"""HTML-роуты пользовательского кабинета.""" + +from fastapi import APIRouter, Depends, Request +from fastapi.templating import Jinja2Templates + +from core.deps import get_current_user_or_none + +templates = Jinja2Templates(directory="templates") +router = APIRouter() + + +@router.get("/login") +async def login_page(request: Request): + """Страница входа.""" + return templates.TemplateResponse("auth/login.html", {"request": request}) + + +@router.get("/register") +async def register_page(request: Request): + """Страница регистрации.""" + return templates.TemplateResponse("auth/register.html", {"request": request}) + + +@router.get("/dashboard") +async def dashboard_page(request: Request, user=Depends(get_current_user_or_none)): + """Главная панель пользователя.""" + return templates.TemplateResponse( + "user/dashboard.html", + {"request": request, "user": user, "access_token": ""}, + ) + + +@router.get("/asr") +async def asr_page(request: Request, user=Depends(get_current_user_or_none)): + """Интерфейс распознавания.""" + return templates.TemplateResponse( + "user/asr.html", + {"request": request, "user": user, "access_token": ""}, + ) + + +@router.get("/history") +async def history_page(request: Request, user=Depends(get_current_user_or_none)): + """История сессий.""" + return templates.TemplateResponse( + "user/history.html", + {"request": request, "user": user, "access_token": ""}, + ) + + +@router.get("/history/{session_id}") +async def history_detail_page(request: Request, session_id: str, user=Depends(get_current_user_or_none)): + """Детали сессии.""" + return templates.TemplateResponse( + "user/history_detail.html", + {"request": request, "user": user, "session_id": session_id, "access_token": ""}, + ) + + +@router.get("/subscription") +async def subscription_page(request: Request, user=Depends(get_current_user_or_none)): + """Управление подпиской.""" + return templates.TemplateResponse( + "user/subscription.html", + {"request": request, "user": user, "access_token": ""}, + ) + + +@router.get("/profile") +async def profile_page(request: Request, user=Depends(get_current_user_or_none)): + """Профиль пользователя.""" + return templates.TemplateResponse( + "user/profile.html", + {"request": request, "user": user, "access_token": ""}, + ) + + +@router.get("/api-keys") +async def api_keys_page(request: Request, user=Depends(get_current_user_or_none)): + """Управление API-ключами.""" + return templates.TemplateResponse( + "user/api_keys.html", + {"request": request, "user": user, "access_token": ""}, + ) diff --git a/services/admin_service.py b/services/admin_service.py new file mode 100644 index 0000000..e19d524 --- /dev/null +++ b/services/admin_service.py @@ -0,0 +1,329 @@ +"""Бизнес-логика админ-панели.""" + +from datetime import datetime, timedelta, timezone +from typing import Any, Optional + +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from core.security import create_access_token # type: ignore[import-untyped] +from db.enums import SubscriptionStatus, UserRole +from db.models import ( + AdminAuditLog, + ApiKey, + ASRSession, + Plan, + Subscription, + SystemLog, + TelegramBotConfig, + Transaction, + User, +) + + +# In-memory флаг режима обслуживания +_maintenance_mode: bool = False + + +def is_maintenance_mode() -> bool: + """Возвращает True, если включён режим обслуживания.""" + return _maintenance_mode + + +def set_maintenance_mode(enabled: bool) -> None: + """Включает/выключает режим обслуживания.""" + global _maintenance_mode + _maintenance_mode = enabled + + +async def get_users_list( + db: AsyncSession, + page: int = 1, + per_page: int = 20, + search: Optional[str] = None, + role: Optional[str] = None, + is_active: Optional[bool] = None, +) -> tuple[list[User], int]: + """Возвращает список пользователей с пагинацией и общее количество.""" + query = select(User) + count_query = select(func.count(User.id)) + + if search: + pattern = f"%{search}%" + query = query.where( + (User.email.ilike(pattern)) | (User.full_name.ilike(pattern)) + ) + count_query = count_query.where( + (User.email.ilike(pattern)) | (User.full_name.ilike(pattern)) + ) + + if role: + query = query.where(User.role == role) + count_query = count_query.where(User.role == role) + + if is_active is not None: + query = query.where(User.is_active.is_(is_active)) + count_query = count_query.where(User.is_active.is_(is_active)) + + query = query.order_by(User.created_at.desc()) + query = query.offset((page - 1) * per_page).limit(per_page) + + result = await db.execute(query) + users = result.scalars().all() + + total_result = await db.execute(count_query) + total = total_result.scalar() or 0 + + return list(users), total + + +async def get_user_by_id(db: AsyncSession, user_id: str) -> Optional[User]: + """Возвращает пользователя по ID.""" + result = await db.execute(select(User).where(User.id == user_id)) + return result.scalar_one_or_none() + + +async def update_user( + db: AsyncSession, + user: User, + data: dict[str, Any], +) -> User: + """Обновляет поля пользователя.""" + for field in ("full_name", "phone", "role", "is_active"): + if field in data and data[field] is not None: + setattr(user, field, data[field]) + await db.commit() + await db.refresh(user) + return user + + +async def get_user_sessions( + db: AsyncSession, + user_id: str, + limit: int = 50, +) -> list[ASRSession]: + """Возвращает ASR-сессии пользователя.""" + result = await db.execute( + select(ASRSession) + .where(ASRSession.user_id == user_id) + .order_by(ASRSession.created_at.desc()) + .limit(limit) + ) + return list(result.scalars().all()) + + +async def impersonate_user(user: User) -> str: + """Генерирует access token от имени пользователя (superadmin only).""" + return create_access_token({"sub": user.id, "role": user.role.value}) + + +async def get_plans(db: AsyncSession) -> list[Plan]: + """Возвращает все тарифные планы.""" + result = await db.execute(select(Plan).order_by(Plan.created_at.desc())) + return list(result.scalars().all()) + + +async def create_plan(db: AsyncSession, data: dict[str, Any]) -> Plan: + """Создаёт новый тарифный план.""" + plan = Plan(**data) + db.add(plan) + await db.commit() + await db.refresh(plan) + return plan + + +async def update_plan( + db: AsyncSession, + plan: Plan, + data: dict[str, Any], +) -> Plan: + """Обновляет тарифный план.""" + for field in ( + "code", + "name", + "description", + "max_requests_per_minute", + "max_audio_duration_sec", + "price_per_month", + "is_active", + ): + if field in data and data[field] is not None: + setattr(plan, field, data[field]) + await db.commit() + await db.refresh(plan) + return plan + + +async def delete_plan(db: AsyncSession, plan: Plan) -> None: + """Деактивирует тарифный план.""" + plan.is_active = False + await db.commit() + + +async def get_subscriptions( + db: AsyncSession, + page: int = 1, + per_page: int = 20, +) -> tuple[list[Subscription], int]: + """Возвращает подписки с пагинацией.""" + result = await db.execute( + select(Subscription) + .order_by(Subscription.created_at.desc()) + .offset((page - 1) * per_page) + .limit(per_page) + ) + total_result = await db.execute(select(func.count(Subscription.id))) + return list(result.scalars().all()), total_result.scalar() or 0 + + +async def extend_subscription( + db: AsyncSession, + subscription: Subscription, + days: int = 30, +) -> Subscription: + """Ручное продление подписки.""" + now = datetime.now(timezone.utc) + if subscription.expires_at: + subscription.expires_at = subscription.expires_at + timedelta(days=days) + else: + subscription.expires_at = now + timedelta(days=days) + subscription.status = SubscriptionStatus.active + await db.commit() + await db.refresh(subscription) + return subscription + + +async def cancel_subscription_admin( + db: AsyncSession, + subscription: Subscription, +) -> None: + """Ручная отмена подписки админом.""" + subscription.status = SubscriptionStatus.cancelled + subscription.auto_renew = False + await db.commit() + + +async def get_transactions( + db: AsyncSession, + page: int = 1, + per_page: int = 20, +) -> tuple[list[Transaction], int]: + """Возвращает транзакции с пагинацией.""" + result = await db.execute( + select(Transaction) + .order_by(Transaction.created_at.desc()) + .offset((page - 1) * per_page) + .limit(per_page) + ) + total_result = await db.execute(select(func.count(Transaction.id))) + return list(result.scalars().all()), total_result.scalar() or 0 + + +async def get_system_logs( + db: AsyncSession, + page: int = 1, + per_page: int = 50, + level: Optional[str] = None, + component: Optional[str] = None, +) -> tuple[list[SystemLog], int]: + """Возвращает системные логи.""" + query = select(SystemLog).order_by(SystemLog.created_at.desc()) + count_query = select(func.count(SystemLog.id)) + + if level: + query = query.where(SystemLog.level == level) + count_query = count_query.where(SystemLog.level == level) + if component: + query = query.where(SystemLog.component == component) + count_query = count_query.where(SystemLog.component == component) + + query = query.offset((page - 1) * per_page).limit(per_page) + result = await db.execute(query) + total_result = await db.execute(count_query) + return list(result.scalars().all()), total_result.scalar() or 0 + + +async def get_audit_logs( + db: AsyncSession, + page: int = 1, + per_page: int = 50, +) -> tuple[list[AdminAuditLog], int]: + """Возвращает аудит-логи админов.""" + result = await db.execute( + select(AdminAuditLog) + .order_by(AdminAuditLog.created_at.desc()) + .offset((page - 1) * per_page) + .limit(per_page) + ) + total_result = await db.execute(select(func.count(AdminAuditLog.id))) + return list(result.scalars().all()), total_result.scalar() or 0 + + +async def get_api_keys( + db: AsyncSession, + page: int = 1, + per_page: int = 50, + user_id: Optional[str] = None, +) -> tuple[list[ApiKey], int]: + """Возвращает API-ключи (все или по пользователю).""" + query = select(ApiKey).order_by(ApiKey.created_at.desc()) + count_query = select(func.count(ApiKey.id)) + + if user_id: + query = query.where(ApiKey.user_id == user_id) + count_query = count_query.where(ApiKey.user_id == user_id) + + query = query.offset((page - 1) * per_page).limit(per_page) + result = await db.execute(query) + total_result = await db.execute(count_query) + return list(result.scalars().all()), total_result.scalar() or 0 + + +async def revoke_api_key_admin(db: AsyncSession, key: ApiKey) -> None: + """Отзывает API-ключ админом.""" + key.is_active = False + await db.commit() + + +async def get_queue_status() -> dict[str, Any]: + """Заглушка: возвращает статус очереди ASR.""" + return {"active_tasks": 0, "pending_tasks": 0} + + +async def cancel_task(task_id: str) -> bool: + """Заглушка: отмена задачи в очереди.""" + return False + + +async def disconnect_user_session(user_id: str, session_id: str) -> bool: + """Заглушка: принудительное отключение WS-сессии.""" + return False + + +async def get_telegram_config(db: AsyncSession) -> Optional[TelegramBotConfig]: + """Возвращает конфигурацию Telegram-бота.""" + result = await db.execute( + select(TelegramBotConfig).where(TelegramBotConfig.is_active.is_(True)) + ) + return result.scalar_one_or_none() + + +async def get_telegram_stats(db: AsyncSession) -> dict[str, int]: + """Возвращает статистику Telegram-пользователей.""" + total_result = await db.execute( + select(func.count(User.id)).where(User.telegram_id.isnot(None)) + ) + return { + "total_telegram_users": total_result.scalar() or 0, + "total_webapp_sessions": 0, + } + + +async def set_telegram_webhook(url: str, bot_token: str) -> dict[str, str]: + """Заглушка: установка webhook Telegram-бота.""" + return {"detail": "Webhook установка — заглушка", "url": url} + + +async def broadcast_message(message: str) -> dict[str, str]: + """Заглушка: рассылка сообщения всем пользователям бота.""" + return {"detail": "Рассылка — заглушка", "message": message} diff --git a/services/asr_pipeline.py b/services/asr_pipeline.py index fd8c10a..c68613c 100644 --- a/services/asr_pipeline.py +++ b/services/asr_pipeline.py @@ -91,6 +91,14 @@ async def process_audio_stream_chunk( # --- 5. Проверка порога VAD --- combined_duration = (session.audio_overlap + session.audio_buffer).duration_seconds + logger.debug( + "Chunk received for %s: chunk=%.3f sec, buffer=%.3f sec, overlap=%.3f sec, combined=%.3f sec", + session.client_id, + audiosegment_chunk.duration_seconds, + session.audio_buffer.duration_seconds, + session.audio_overlap.duration_seconds, + combined_duration, + ) if combined_duration < settings.MAX_OVERLAP_DURATION: return @@ -111,6 +119,19 @@ async def process_audio_stream_chunk( pass return + # Fallback: если VAD не перенёс ничего в audio_to_asr, принудительно сливаем весь буфер + if not session.audio_to_asr: + logger.warning( + "VAD produced empty audio_to_asr for %s (buffer=%.3f sec, overlap=%.3f sec). " + "Forcing entire buffer to audio_to_asr.", + session.client_id, + session.audio_buffer.duration_seconds, + session.audio_overlap.duration_seconds, + ) + session.audio_to_asr.append(session.audio_buffer) + session.audio_overlap = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + session.audio_buffer = AudioSegment.silent(1, frame_rate=settings.BASE_SAMPLE_RATE) + # --- 6. Распознавание последнего сегмента --- if not session.audio_to_asr: return @@ -132,6 +153,14 @@ async def process_audio_stream_chunk( if metrics_collector is not None: metrics_collector.decrement_tasks() + logger.debug( + "Chunk recognized for %s: segment=%.3f sec, audio_duration=%.3f sec, text='%s'", + session.client_id, + segment.duration_seconds, + session.audio_duration, + asr_result_words.get("data", {}).get("text", "")[:50], + ) + # --- 7. Отправка результата --- text = asr_result_words.get("data", {}).get("text", "") is_silence = len(text) == 0 or text == " " @@ -206,8 +235,15 @@ async def process_final_audio( try: # --- 1. Объединение остатков --- final_audio = session.audio_overlap + session.audio_buffer - session.audio_to_asr.append(final_audio) - logger.debug("Final audio duration: %.3f sec", final_audio.duration_seconds) + if final_audio.duration_seconds > 0.1: + session.audio_to_asr.append(final_audio) + logger.debug( + "Final audio for %s: duration=%.3f sec, audio_to_asr count=%d, audio_duration=%.3f sec", + session.client_id, + final_audio.duration_seconds, + len(session.audio_to_asr), + session.audio_duration, + ) # --- 2. Дополнение тишиной при необходимости --- if final_audio.duration_seconds < 2: @@ -215,13 +251,19 @@ async def process_final_audio( session.audio_to_asr[-1] = final_audio logger.debug("Final audio padded with silence to %.3f sec", final_audio.duration_seconds) - # --- 3. Распознавание --- + # --- 3. Распознавание финального остатка --- if metrics_collector is not None: metrics_collector.increment_tasks() try: last_asr_result = await simple_recognise(final_audio, recognizer=recognizer) last_result = process_single_token_vocab_output(last_asr_result, session.audio_duration) session.ws_collected_asr_res[f"channel_{1}"].append(last_result) + logger.debug( + "Final chunk recognized for %s: final_duration=%.3f sec, time_shift=%.3f sec", + session.client_id, + final_audio.duration_seconds, + session.audio_duration, + ) finally: if metrics_collector is not None: metrics_collector.decrement_tasks() diff --git a/services/auth_service.py b/services/auth_service.py new file mode 100644 index 0000000..7e0c2d2 --- /dev/null +++ b/services/auth_service.py @@ -0,0 +1,146 @@ +"""Бизнес-логика аутентификации (JWT, Telegram, регистрация).""" + +import hashlib +import hmac +import time +import urllib.parse +from datetime import datetime, timezone + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from db.enums import UserRole +from db.models import User +from core.security import ( # type: ignore[import-untyped] + create_access_token, + create_refresh_token, + get_password_hash, + verify_password, +) + + +# In-memory blacklist для refresh-токенов (в production → Redis) +_refresh_blacklist: dict[str, float] = {} + + +async def is_refresh_blacklisted(token: str) -> bool: + """Проверяет, не отозван ли refresh-токен.""" + exp = _refresh_blacklist.get(token) + if not exp: + return False + if time.time() > exp: + _refresh_blacklist.pop(token, None) + return False + return True + + +async def blacklist_refresh_token(token: str, ttl_sec: int = 7 * 24 * 3600) -> None: + """Добавляет refresh-токен в чёрный список.""" + _refresh_blacklist[token] = time.time() + ttl_sec + + +async def authenticate_user(db: AsyncSession, email: str, password: str) -> User | None: + """Проверяет email/пароль и возвращает пользователя.""" + result = await db.execute(select(User).where(User.email == email)) + user: User | None = result.scalar_one_or_none() + if not user or not user.is_active: + return None + if not verify_password(password, user.hashed_password): + return None + return user + + +def generate_tokens(user: User) -> dict[str, str]: + """Генерирует пару access/refresh токенов.""" + role_str = user.role.value if hasattr(user.role, "value") else user.role + access = create_access_token({"sub": user.id, "role": role_str}) + refresh = create_refresh_token({"sub": user.id}) + return {"access_token": access, "refresh_token": refresh} + + +async def register_user( + db: AsyncSession, email: str, password: str, full_name: str | None = None +) -> User: + """Регистрирует нового пользователя.""" + result = await db.execute(select(User).where(User.email == email)) + if result.scalar_one_or_none(): + raise ValueError("Пользователь с таким email уже существует") + + user = User( + email=email, + hashed_password=get_password_hash(password), + full_name=full_name, + role=UserRole.user, + ) + db.add(user) + await db.commit() + await db.refresh(user) + return user + + +def validate_telegram_init_data(init_data: str, bot_token: str) -> dict | None: + """Проверяет подпись initData от Telegram Web App.""" + try: + parsed = dict(urllib.parse.parse_qsl(init_data)) + received_hash = parsed.pop("hash", None) + if not received_hash: + return None + + data_check_string = "\n".join(f"{k}={v}" for k, v in sorted(parsed.items())) + + secret_key = hmac.new( + key=b"WebAppData", + msg=bot_token.encode(), + digestmod=hashlib.sha256, + ).digest() + + calculated_hash = hmac.new( + key=secret_key, + msg=data_check_string.encode(), + digestmod=hashlib.sha256, + ).hexdigest() + + if not hmac.compare_digest(calculated_hash, received_hash): + return None + + return parsed + except Exception: + return None + + +async def get_or_create_telegram_user(db: AsyncSession, tg_data: dict) -> User: + """Находит или создаёт пользователя по telegram_id.""" + telegram_id = int(tg_data.get("id", 0)) + if not telegram_id: + raise ValueError("Отсутствует telegram_id") + + result = await db.execute(select(User).where(User.telegram_id == telegram_id)) + user: User | None = result.scalar_one_or_none() + + if user: + # Обновляем Telegram-поля + user.telegram_username = tg_data.get("username") or user.telegram_username + user.telegram_first_name = tg_data.get("first_name") or user.telegram_first_name + user.telegram_last_name = tg_data.get("last_name") or user.telegram_last_name + user.telegram_photo_url = tg_data.get("photo_url") or user.telegram_photo_url + user.telegram_auth_date = datetime.now(timezone.utc) + await db.commit() + await db.refresh(user) + return user + + # Создаём нового пользователя без пароля (только Telegram) + user = User( + email=f"tg_{telegram_id}@telegram.local", + hashed_password="", + telegram_id=telegram_id, + telegram_username=tg_data.get("username"), + telegram_first_name=tg_data.get("first_name"), + telegram_last_name=tg_data.get("last_name"), + telegram_photo_url=tg_data.get("photo_url"), + telegram_auth_date=datetime.now(timezone.utc), + role=UserRole.user, + ) + db.add(user) + await db.commit() + await db.refresh(user) + return user diff --git a/services/metrics_reporter.py b/services/metrics_reporter.py new file mode 100644 index 0000000..704c3b8 --- /dev/null +++ b/services/metrics_reporter.py @@ -0,0 +1,50 @@ +"""Фоновая задача записи системных метрик в БД.""" + +import asyncio +import logging + +from db.enums import SystemLogLevel +from db.models import SystemLog +from db.session import AsyncSessionLocal + +logger = logging.getLogger(__name__) + + +async def metrics_reporter_loop(app_state, interval_sec: float = 30.0): + """Каждые interval_sec секунд собирает метрики и пишет в SystemLog.""" + while True: + try: + await asyncio.sleep(interval_sec) + + metrics_collector = getattr(app_state, "metrics_collector", None) + ws_manager = getattr(app_state, "ws_manager", None) + + if not metrics_collector: + continue + + active_conn = ws_manager.active_connections_count if ws_manager else 0 + max_conn = getattr(ws_manager, "max_connections", 100) + + metrics = metrics_collector.collect( + active_connections=active_conn, + max_connections=max_conn, + ) + + # Преобразуем Pydantic-модель WSStatusResponse в dict для JSONB + metrics_dict = metrics.model_dump() if hasattr(metrics, "model_dump") else metrics.dict() + + async with AsyncSessionLocal() as session: + log = SystemLog( + level=SystemLogLevel.info, + component="SystemMetricsCollector", + message="Periodic metrics snapshot", + meta=metrics_dict, + ) + session.add(log) + await session.commit() + logger.debug("Метрики записаны в SystemLog") + + except asyncio.CancelledError: + break + except Exception as exc: + logger.error(f"Ошибка в metrics_reporter_loop: {exc}") diff --git a/services/ws_manager.py b/services/ws_manager.py index 4af8df6..b8d0ad2 100644 --- a/services/ws_manager.py +++ b/services/ws_manager.py @@ -38,6 +38,7 @@ class ConnectionMeta: client_ip: Optional[str] = None user_agent: Optional[str] = None subscribe_status: bool = False + user_id: Optional[str] = None class ConnectionManager: diff --git a/services/ws_session.py b/services/ws_session.py index e2d8729..4d1cb54 100644 --- a/services/ws_session.py +++ b/services/ws_session.py @@ -43,6 +43,7 @@ def __init__( """ super().__init__(session_id=client_id) self.state = SessionState.connecting + self.user_id: Optional[str] = None self.wait_null_answers: bool = True self.last_activity: float = time.time() self.max_buffer_duration_sec: float = ( diff --git a/static/css/design-system.css b/static/css/design-system.css new file mode 100644 index 0000000..479127b --- /dev/null +++ b/static/css/design-system.css @@ -0,0 +1,147 @@ +:root { + --color-bg: #0f172a; + --color-surface: #1e293b; + --color-primary: #3b82f6; + --color-text: #e2e8f0; + --color-muted: #94a3b8; + --color-success: #22c55e; + --color-warning: #f59e0b; + --color-danger: #ef4444; + --color-border: #334155; + --radius: 8px; + --spacing: 16px; +} + +* { box-sizing: border-box; } +body { + margin: 0; + font-family: system-ui, -apple-system, Segoe UI, Roboto, sans-serif; + background: var(--color-bg); + color: var(--color-text); + line-height: 1.5; +} + +.card { + background: var(--color-surface); + border: 1px solid var(--color-border); + border-radius: var(--radius); + padding: var(--spacing); +} + +.btn { + display: inline-flex; + align-items: center; + justify-content: center; + gap: 8px; + padding: 10px 16px; + border: 1px solid transparent; + border-radius: var(--radius); + font-size: 14px; + font-weight: 500; + cursor: pointer; + transition: opacity .2s, transform .1s; +} +.btn:disabled { opacity: .6; cursor: not-allowed; } +.btn-primary { background: var(--color-primary); color: #fff; } +.btn-danger { background: var(--color-danger); color: #fff; } +.btn-ghost { background: transparent; border-color: var(--color-border); color: var(--color-text); } + +.form-group { margin-bottom: 12px; } +.form-group label { display: block; margin-bottom: 4px; font-size: 13px; color: var(--color-muted); } +.input { + width: 100%; + padding: 10px 12px; + background: var(--color-bg); + border: 1px solid var(--color-border); + border-radius: var(--radius); + color: var(--color-text); + font-size: 14px; +} +.input:focus { outline: none; border-color: var(--color-primary); } + +.badge { + display: inline-block; + padding: 2px 8px; + border-radius: 999px; + font-size: 12px; + font-weight: 600; +} +.badge-success { background: rgba(34,197,94,.15); color: var(--color-success); } +.badge-warning { background: rgba(245,158,11,.15); color: var(--color-warning); } +.badge-danger { background: rgba(239,68,68,.15); color: var(--color-danger); } + +.skeleton { + background: linear-gradient(90deg, var(--color-surface) 25%, #334155 50%, var(--color-surface) 75%); + background-size: 200% 100%; + animation: skeleton 1.2s infinite; + border-radius: var(--radius); +} +@keyframes skeleton { + 0% { background-position: 200% 0; } + 100% { background-position: -200% 0; } +} + +.toast-container { + position: fixed; + top: 16px; + right: 16px; + z-index: 9999; + display: flex; + flex-direction: column; + gap: 8px; + max-width: 360px; +} +.toast { + padding: 12px 16px; + border-radius: var(--radius); + background: var(--color-surface); + border: 1px solid var(--color-border); + color: var(--color-text); + font-size: 14px; + box-shadow: 0 10px 30px rgba(0,0,0,.3); + animation: toastIn .3s ease; +} +.toast.info { border-left: 4px solid var(--color-primary); } +.toast.success { border-left: 4px solid var(--color-success); } +.toast.warning { border-left: 4px solid var(--color-warning); } +.toast.error { border-left: 4px solid var(--color-danger); } +.toast .progress { + height: 3px; + background: var(--color-primary); + margin-top: 8px; + border-radius: 2px; + animation: progress linear; +} +@keyframes toastIn { + from { transform: translateX(100%); opacity: 0; } + to { transform: translateX(0); opacity: 1; } +} +@keyframes progress { + from { width: 100%; } + to { width: 0%; } +} + +.spinner { + width: 16px; + height: 16px; + border: 2px solid var(--color-border); + border-top-color: var(--color-primary); + border-radius: 50%; + animation: spin .8s linear infinite; +} +@keyframes spin { to { transform: rotate(360deg); } } + +.container { max-width: 1200px; margin: 0 auto; padding: 0 var(--spacing); } +.flex { display: flex; } +.flex-col { flex-direction: column; } +.items-center { align-items: center; } +.justify-between { justify-content: space-between; } +.gap-2 { gap: 8px; } +.gap-4 { gap: 16px; } +.mt-4 { margin-top: 16px; } +.mb-4 { margin-bottom: 16px; } +.w-full { width: 100%; } +.h-full { height: 100%; } +.text-sm { font-size: 14px; } +.text-xs { font-size: 12px; } +.text-muted { color: var(--color-muted); } diff --git a/static/js/asr_client.js b/static/js/asr_client.js new file mode 100644 index 0000000..f24020e --- /dev/null +++ b/static/js/asr_client.js @@ -0,0 +1,113 @@ +(function() { + class ASRClient { + constructor() { + this.ws = null; + this.reconnectAttempts = 0; + this.maxReconnectAttempts = 5; + this.baseDelay = 1000; + this.maxDelay = 30000; + this._onPartial = null; + this._onFinal = null; + this._onError = null; + this._pendingMessages = []; + this._isConnected = false; + this._authToken = null; + this._config = null; + } + + onPartial(cb) { this._onPartial = cb; } + onFinal(cb) { this._onFinal = cb; } + onError(cb) { this._onError = cb; } + + connect(authToken, config) { + this._authToken = authToken; + this._config = config; + this._connect(); + } + + _connect() { + if (this.ws) { + try { this.ws.close(); } catch(e) {} + } + + const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; + const wsUrl = protocol + '//' + window.location.host + '/api/v1/asr/ws'; + + this.ws = new WebSocket(wsUrl); + this.ws.binaryType = 'arraybuffer'; + + this.ws.onopen = () => { + this.reconnectAttempts = 0; + this._isConnected = true; + if (this._authToken) { + this.ws.send(JSON.stringify({ type: 'auth', access_token: this._authToken })); + } + this.ws.send(JSON.stringify({ type: 'config', ...this._config })); + while (this._pendingMessages.length > 0) { + const msg = this._pendingMessages.shift(); + this.ws.send(msg); + } + }; + + this.ws.onmessage = (event) => { + try { + const msg = JSON.parse(event.data); + if (msg.type === 'partial_result' && this._onPartial) { + this._onPartial(msg); + } else if (msg.type === 'final_result' && this._onFinal) { + this._onFinal(msg); + } else if (msg.type === 'error' && this._onError) { + this._onError(msg); + } + } catch (e) { + console.error('WS parse error', e); + } + }; + + this.ws.onerror = () => { + if (this._onError) this._onError({ code: 'ws_error', message: 'WebSocket error' }); + }; + + this.ws.onclose = (event) => { + this._isConnected = false; + if (!event.wasClean && this.reconnectAttempts < this.maxReconnectAttempts) { + this._doReconnect(); + } + }; + } + + _doReconnect() { + const delay = Math.min(this.baseDelay * Math.pow(2, this.reconnectAttempts), this.maxDelay); + this.reconnectAttempts++; + setTimeout(() => this._connect(), delay); + } + + sendAudioChunk(base64Chunk, seqNum) { + const payload = JSON.stringify({ type: 'audio_chunk', audio_base64: base64Chunk, seq_num: seqNum }); + if (this._isConnected && this.ws && this.ws.readyState === WebSocket.OPEN) { + this.ws.send(payload); + } else { + this._pendingMessages.push(payload); + } + } + + sendEOS() { + const payload = JSON.stringify({ type: 'eos' }); + if (this._isConnected && this.ws && this.ws.readyState === WebSocket.OPEN) { + this.ws.send(payload); + } else { + this._pendingMessages.push(payload); + } + } + + disconnect() { + this.reconnectAttempts = this.maxReconnectAttempts; + if (this.ws) { + try { this.ws.close(1000, 'Client disconnect'); } catch(e) {} + } + this._isConnected = false; + } + } + + window.ASRClient = ASRClient; +})(); diff --git a/static/js/asr_page.js b/static/js/asr_page.js new file mode 100644 index 0000000..23a9a76 --- /dev/null +++ b/static/js/asr_page.js @@ -0,0 +1,296 @@ +(function() { + 'use strict'; + + // --- Табы --- + function switchTab(name) { + document.querySelectorAll('.tab-panel').forEach(p => p.style.display = 'none'); + document.querySelectorAll('.tab-btn').forEach(b => b.classList.remove('active')); + document.getElementById('panel-' + name).style.display = 'block'; + document.getElementById('tab-' + name).classList.add('active'); + } + window.switchTab = switchTab; + + // --- Утилиты --- + function copyText(elId) { + const text = document.getElementById(elId).textContent; + navigator.clipboard.writeText(text).then(() => UI.toast('Скопировано', 'success')); + } + window.copyText = copyText; + + function downloadText(elId, filename) { + const text = document.getElementById(elId).textContent; + const blob = new Blob([text], {type:'text/plain'}); + const a = document.createElement('a'); + a.href = URL.createObjectURL(blob); + a.download = filename; + a.click(); + } + window.downloadText = downloadText; + + // --- URL --- + async function sendUrl() { + const btn = document.getElementById('btnUrl'); + UI.setLoading(btn, true); + const url = document.getElementById('urlInput').value.trim(); + if (!url) { UI.toast('Введите ссылку', 'warning'); UI.setLoading(btn, false); return; } + try { + const resp = await Auth.apiFetch('/api/v1/asr/url', { + method: 'POST', + headers: {'Content-Type': 'application/json'}, + body: JSON.stringify({ + AudioFileUrl: url, + keep_raw: document.getElementById('url_keep_raw').checked, + do_echo_clearing: document.getElementById('url_do_echo').checked, + do_dialogue: document.getElementById('url_do_dialogue').checked, + do_punctuation: document.getElementById('url_do_punct').checked + }) + }); + const data = await resp.json(); + const el = document.getElementById('urlResult'); + const pre = document.getElementById('urlResultText'); + el.style.display = 'block'; + pre.textContent = data.data?.sentenced_data?.raw_text_sentenced_recognition || data.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || JSON.stringify(data, null, 2); + } catch (e) { + UI.toast('Ошибка: ' + e.message, 'error'); + } finally { + UI.setLoading(btn, false); + } + } + window.sendUrl = sendUrl; + + // --- Файл (drag & drop) --- + let selectedFile = null; + + function handleDrop(e) { + e.preventDefault(); + e.currentTarget.style.borderColor = 'var(--color-border)'; + if (e.dataTransfer.files.length) handleFileSelect(e.dataTransfer.files[0]); + } + window.handleDrop = handleDrop; + + function handleFileSelect(file) { + selectedFile = file; + document.getElementById('btnFile').disabled = false; + UI.toast('Файл выбран: ' + file.name, 'info'); + } + window.handleFileSelect = handleFileSelect; + + async function sendFile() { + if (!selectedFile) return; + const btn = document.getElementById('btnFile'); + UI.setLoading(btn, true); + document.getElementById('fileProgress').style.display = 'block'; + const form = new FormData(); + form.append('file', selectedFile); + form.append('keep_raw', document.getElementById('file_keep_raw').checked); + form.append('do_dialogue', document.getElementById('file_do_dialogue').checked); + form.append('do_diarization', document.getElementById('file_do_diar').checked); + form.append('do_punctuation', document.getElementById('file_do_punct').checked); + try { + const resp = await Auth.apiFetch('/api/v1/asr/file', {method:'POST', body:form}); + const data = await resp.json(); + document.getElementById('fileResult').style.display = 'block'; + document.getElementById('fileResultText').textContent = data.data?.sentenced_data?.raw_text_sentenced_recognition || data.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || JSON.stringify(data, null, 2); + } catch (e) { + UI.toast('Ошибка: ' + e.message, 'error'); + } finally { + UI.setLoading(btn, false); + document.getElementById('fileProgress').style.display = 'none'; + } + } + window.sendFile = sendFile; + + document.getElementById('dropZone').addEventListener('click', () => document.getElementById('fileInput').click()); + + // --- WebSocket (legacy logic from index.html, adapted for new protocol) --- + let wsSockets = []; + + function onWsFileSelected() { + const f = document.getElementById('wsFileInput').files[0]; + document.getElementById('btnWs').disabled = !f; + } + window.onWsFileSelected = onWsFileSelected; + + function setWsStatus(status) { + const dot = document.getElementById('wsStatusDot'); + const txt = document.getElementById('wsStatusText'); + if (status === 'connected') { dot.style.background = 'var(--color-success)'; txt.textContent = 'Подключено'; } + else if (status === 'connecting') { dot.style.background = 'var(--color-warning)'; txt.textContent = 'Подключение...'; } + else { dot.style.background = 'var(--color-danger)'; txt.textContent = 'Отключено'; } + } + + function readWavHeader(arrayBuffer) { + const dataView = new DataView(arrayBuffer); + if (String.fromCharCode(...new Uint8Array(arrayBuffer, 0, 4)) !== 'RIFF') { + throw new Error('Файл не является WAV файлом'); + } + const sampleRate = dataView.getUint32(24, true); + const numChannels = dataView.getUint16(22, true); + const bitsPerSample = dataView.getUint16(34, true); + return { sampleRate, numChannels, bitsPerSample }; + } + + function arrayBufferToBase64(buffer) { + let binary = ''; + const bytes = new Uint8Array(buffer); + const len = bytes.byteLength; + for (let i = 0; i < len; i++) { + binary += String.fromCharCode(bytes[i]); + } + return window.btoa(binary); + } + + async function sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate) { + return new Promise((resolve, reject) => { + const wsProtocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; + const wsUrl = `${wsProtocol}//${window.location.host}/api/v1/asr/ws`; + const socket = new WebSocket(wsUrl); + wsSockets.push(socket); + + const doDialogue = document.getElementById('ws_do_dialogue').checked; + const doPunctuation = document.getElementById('ws_do_punct').checked; + const token = Auth.getAccessToken(); + + socket.onopen = async function() { + const useBase64 = document.getElementById('ws_use_base64').checked; + if (token) { + socket.send(JSON.stringify({ type: 'auth', access_token: token })); + } + socket.send(JSON.stringify({ + type: 'config', + sample_rate: sampleRate, + audio_format: 'pcm16', + audio_transport: useBase64 ? 'json_base64' : 'binary', + wait_null_answers: false, + do_dialogue: doDialogue, + do_punctuation: doPunctuation, + channel_name: 'channel_' + (channel + 1) + })); + + const dataView = new Uint8Array(arrayBuffer); + const bytesPerSample = 2; + let offset = 44 + channel * bytesPerSample; + let seqNum = 0; + + while (offset < dataView.length) { + const endOffset = Math.min(offset + chunkSize * numChannels, dataView.length); + const chunk = new Uint8Array((endOffset - offset) / numChannels); + + for (let i = 0; i < chunk.length; i += bytesPerSample) { + const sampleOffset = offset + i * numChannels; + chunk[i] = dataView[sampleOffset]; + chunk[i + 1] = dataView[sampleOffset + 1]; + } + + if (useBase64) { + const base64Chunk = arrayBufferToBase64(chunk); + socket.send(JSON.stringify({ + type: 'audio_chunk', + audio_base64: base64Chunk, + seq_num: seqNum++ + })); + } else { + socket.send(chunk); + } + offset += chunkSize * numChannels; + await new Promise((resolve) => setTimeout(resolve, 100)); + } + + socket.send(JSON.stringify({ type: 'eos' })); + }; + + socket.onmessage = function(event) { + const data = JSON.parse(event.data); + if (data.type === 'partial_result') { + const text = data.data?.text || ''; + document.getElementById('wsPartial').textContent = `[Канал ${channel+1}]: ${text} ...`; + } + if (data.type === 'final_result' || data.last_message) { + const text = data.sentenced_data?.raw_text_sentenced_recognition || data.data?.text || JSON.stringify(data, null, 2); + const current = document.getElementById('wsResultText').textContent; + document.getElementById('wsResultText').textContent = current + (current ? '\n\n' : '') + `=== Канал ${channel+1} ===\n${text}`; + document.getElementById('wsPartial').textContent = ''; + } + if (data.last_message) { + socket.close(); + resolve(); + } + }; + + socket.onerror = function(error) { + UI.toast('Ошибка WebSocket канала ' + (channel+1), 'error'); + reject(error); + }; + + socket.onclose = function() { + console.log(`Сокет для канала ${channel+1} закрыт`); + }; + }); + } + + async function sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate) { + for (let channel = 0; channel < numChannels; channel++) { + await sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate); + console.log("конец канала"); + } + } + + async function sendWs() { + const file = document.getElementById('wsFileInput').files[0]; + if (!file) return; + UI.setLoading(document.getElementById('btnWs'), true); + document.getElementById('wsResult').style.display = 'block'; + document.getElementById('wsResultText').textContent = ''; + document.getElementById('wsPartial').textContent = ''; + document.getElementById('btnWsStop').style.display = 'inline-flex'; + setWsStatus('connecting'); + wsSockets = []; + + try { + const arrayBuffer = await file.arrayBuffer(); + const { sampleRate, numChannels } = readWavHeader(arrayBuffer); + const chunkSize = 65536; + await sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate); + setWsStatus('disconnected'); + UI.toast('Распознавание завершено', 'success'); + } catch (e) { + UI.toast('Ошибка: ' + e.message, 'error'); + setWsStatus('disconnected'); + } finally { + UI.setLoading(document.getElementById('btnWs'), false); + document.getElementById('btnWsStop').style.display = 'none'; + } + } + window.sendWs = sendWs; + + function stopWs() { + wsSockets.forEach(s => { try { s.close(); } catch(e) {} }); + wsSockets = []; + setWsStatus('disconnected'); + UI.setLoading(document.getElementById('btnWs'), false); + document.getElementById('btnWsStop').style.display = 'none'; + } + window.stopWs = stopWs; + + // --- История --- + async function loadHistory() { + const el = document.getElementById('asrHistory'); + try { + const resp = await Auth.apiFetch('/api/v1/user/sessions?limit=5'); + const data = await resp.json(); + if (!data || !data.length) { el.innerHTML = '
Нет сессий
'; return; } + el.innerHTML = '' + + data.map(s => ` + + + + + `).join('') + '
${UI.formatDate(s.created_at)}${s.status}${s.session_type}Открыть
'; + } catch (e) { + el.innerHTML = '
Не удалось загрузить
'; + } + } + + // Инициализация + loadHistory(); +})(); diff --git a/static/js/auth.js b/static/js/auth.js new file mode 100644 index 0000000..3822878 --- /dev/null +++ b/static/js/auth.js @@ -0,0 +1,89 @@ +(function() { + let _token = null; + + function setAccessToken(t) { _token = t; } + function getAccessToken() { return _token; } + function clearAuth() { _token = null; } + + let _isRefreshing = false; + let _refreshPromise = null; + + async function apiFetch(url, opts = {}) { + opts.headers = opts.headers || {}; + if (!opts.credentials) { + opts.credentials = 'include'; + } + + // Если токена нет, пробуем восстановить сессию через refresh cookie + if (!_token) { + await refreshToken(); + } + + if (_token) { + opts.headers['Authorization'] = 'Bearer ' + _token; + } + + let resp = await fetch(url, opts); + if (resp.status === 401) { + const refreshed = await refreshToken(); + if (refreshed) { + opts.headers['Authorization'] = 'Bearer ' + _token; + resp = await fetch(url, opts); + } else { + clearAuth(); + window.location.href = '/login'; + return Promise.reject(new Error('Session expired')); + } + } + if (resp.status === 403) { + clearAuth(); + window.location.href = '/login'; + return Promise.reject(new Error('Forbidden')); + } + return resp; + } + + async function refreshToken() { + if (_isRefreshing) { + return await _refreshPromise; + } + _isRefreshing = true; + _refreshPromise = _doRefresh(); + const result = await _refreshPromise; + _isRefreshing = false; + _refreshPromise = null; + return result; + } + + async function _doRefresh() { + try { + const resp = await fetch('/api/v1/auth/refresh', { + method: 'POST', + credentials: 'include' + }); + if (!resp.ok) return false; + const data = await resp.json(); + if (data.access_token) { + setAccessToken(data.access_token); + return true; + } + return false; + } catch (e) { + return false; + } + } + + async function initAuth() { + const publicPaths = ['/login', '/register', '/tg']; + if (publicPaths.some(p => window.location.pathname.startsWith(p))) return; + if (!_token) { + // Пытаемся восстановить сессию через refresh cookie + const refreshed = await refreshToken(); + if (!refreshed) { + window.location.href = '/login'; + } + } + } + + window.Auth = { setAccessToken, getAccessToken, clearAuth, apiFetch, initAuth }; +})(); diff --git a/static/js/ui.js b/static/js/ui.js new file mode 100644 index 0000000..49f4d76 --- /dev/null +++ b/static/js/ui.js @@ -0,0 +1,83 @@ +(function() { + const container = document.createElement('div'); + container.className = 'toast-container'; + document.body.appendChild(container); + + function toast(msg, type = 'info', duration = 5000) { + const el = document.createElement('div'); + el.className = 'toast ' + type; + el.textContent = msg; + + const progress = document.createElement('div'); + progress.className = 'progress'; + progress.style.animationDuration = duration + 'ms'; + el.appendChild(progress); + + container.appendChild(el); + setTimeout(() => { + el.style.opacity = '0'; + setTimeout(() => el.remove(), 300); + }, duration); + } + + function confirmDialog(title, text, onConfirm) { + const dialog = document.createElement('dialog'); + dialog.innerHTML = ` +
+

${escapeHtml(title)}

+

${escapeHtml(text)}

+
+ + +
+
+ `; + document.body.appendChild(dialog); + dialog.showModal(); + + dialog.querySelector('#dlg-cancel').onclick = () => { dialog.close(); dialog.remove(); }; + dialog.querySelector('#dlg-confirm').onclick = () => { dialog.close(); dialog.remove(); onConfirm(); }; + dialog.addEventListener('close', () => dialog.remove()); + } + + function setLoading(element, isLoading) { + if (!element) return; + if (isLoading) { + element.disabled = true; + element.dataset.originalText = element.innerHTML; + element.innerHTML = ''; + } else { + element.disabled = false; + element.innerHTML = element.dataset.originalText || element.innerHTML; + } + } + + function formatDate(iso) { + if (!iso) return '-'; + const d = new Date(iso); + return d.toLocaleString('ru-RU'); + } + + function formatBytes(bytes) { + if (!bytes) return '0 B'; + const k = 1024; + const sizes = ['B','KB','MB','GB']; + const i = Math.floor(Math.log(bytes) / Math.log(k)); + return parseFloat((bytes / Math.pow(k, i)).toFixed(1)) + ' ' + sizes[i]; + } + + function formatDuration(sec) { + if (!sec) return '0:00'; + const m = Math.floor(sec / 60); + const s = Math.floor(sec % 60); + return m + ':' + String(s).padStart(2, '0'); + } + + function escapeHtml(str) { + const div = document.createElement('div'); + div.textContent = str; + return div.innerHTML; + } + + window.UI = { toast, confirmDialog, setLoading, formatDate, formatBytes, formatDuration }; +})(); diff --git a/static/monitor.html b/static/monitor.html new file mode 100644 index 0000000..39faf7e --- /dev/null +++ b/static/monitor.html @@ -0,0 +1,291 @@ + + + + + + ASR Monitor + + + +

ASR Adapter Monitor

+
+
+ Отключено +
+ +
+
+
Adapter Status
+
+
+
+
Active Tasks
+
+
+
+
Connections
+
+
+
+
Uptime
+
+
+
+
GPU Memory
+
+
+
+
+
GPU Utilization
+
+
+
+
+
CPU Memory
+
+
+
+
+
CPU Utilization
+
+
+
+
+
Queue Depth
+
+
+
+
Temperature
+
+
+
+ +
Event Log (last 20)
+
+ + + + diff --git a/templates/admin/api_keys.html b/templates/admin/api_keys.html new file mode 100644 index 0000000..a9bc578 --- /dev/null +++ b/templates/admin/api_keys.html @@ -0,0 +1,32 @@ +{% extends "admin/base_admin.html" %} +{% block title %}API-ключи{% endblock %} +{% block content %} +

API-ключи

+
+

Загрузка ключей...

+ + + + + +
IDПользовательНазваниеАктивенДействия
+
+ +{% endblock %} diff --git a/templates/admin/base_admin.html b/templates/admin/base_admin.html new file mode 100644 index 0000000..ee7727a --- /dev/null +++ b/templates/admin/base_admin.html @@ -0,0 +1,35 @@ + + + + + + {% block title %}Админ-панель{% endblock %} + + + + + +
+ {% block content %}{% endblock %} +
+ + diff --git a/templates/admin/dashboard.html b/templates/admin/dashboard.html new file mode 100644 index 0000000..13fb0df --- /dev/null +++ b/templates/admin/dashboard.html @@ -0,0 +1,18 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Dashboard{% endblock %} +{% block content %} +

Dashboard

+
+

Загрузка метрик...

+ +
+ +{% endblock %} diff --git a/templates/admin/login.html b/templates/admin/login.html new file mode 100644 index 0000000..118e78c --- /dev/null +++ b/templates/admin/login.html @@ -0,0 +1,41 @@ + + + + + + Вход в админ-панель + + + +
+

Admin Login

+ + + +
+ + + diff --git a/templates/admin/logs.html b/templates/admin/logs.html new file mode 100644 index 0000000..c36c349 --- /dev/null +++ b/templates/admin/logs.html @@ -0,0 +1,31 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Логи{% endblock %} +{% block content %} +

Системные логи

+
+

Загрузка логов...

+ + + + + +
УровеньКомпонентСообщениеДата
+
+ +{% endblock %} diff --git a/templates/admin/sessions.html b/templates/admin/sessions.html new file mode 100644 index 0000000..fe726e5 --- /dev/null +++ b/templates/admin/sessions.html @@ -0,0 +1,8 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Сессии{% endblock %} +{% block content %} +

Активные сессии

+
+

Мониторинг ASR-сессий — заглушка.

+
+{% endblock %} diff --git a/templates/admin/settings.html b/templates/admin/settings.html new file mode 100644 index 0000000..419da3c --- /dev/null +++ b/templates/admin/settings.html @@ -0,0 +1,21 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Настройки{% endblock %} +{% block content %} +

Настройки системы

+
+

Режим обслуживания

+ +

Статус: неизвестно

+
+ +{% endblock %} diff --git a/templates/admin/subscriptions.html b/templates/admin/subscriptions.html new file mode 100644 index 0000000..17d5001 --- /dev/null +++ b/templates/admin/subscriptions.html @@ -0,0 +1,32 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Подписки{% endblock %} +{% block content %} +

Подписки

+
+

Загрузка подписок...

+ + + + + +
IDПользовательПланСтатусДействия
+
+ +{% endblock %} diff --git a/templates/admin/tariffs.html b/templates/admin/tariffs.html new file mode 100644 index 0000000..b12d2d0 --- /dev/null +++ b/templates/admin/tariffs.html @@ -0,0 +1,32 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Тарифы{% endblock %} +{% block content %} +

Тарифные планы

+
+

Загрузка тарифов...

+ + + + + +
КодНазваниеЗапросов/минЦена/месАктивен
+
+ +{% endblock %} diff --git a/templates/admin/telegram.html b/templates/admin/telegram.html new file mode 100644 index 0000000..ae6d604 --- /dev/null +++ b/templates/admin/telegram.html @@ -0,0 +1,21 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Telegram{% endblock %} +{% block content %} +

Telegram-интеграция

+
+

Загрузка конфигурации...

+
+
+ +{% endblock %} diff --git a/templates/admin/transactions.html b/templates/admin/transactions.html new file mode 100644 index 0000000..e917c70 --- /dev/null +++ b/templates/admin/transactions.html @@ -0,0 +1,33 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Транзакции{% endblock %} +{% block content %} +

Транзакции

+
+

Загрузка транзакций...

+ + + + + +
IDПользовательСуммаВалютаСтатусПровайдер
+
+ +{% endblock %} diff --git a/templates/admin/users.html b/templates/admin/users.html new file mode 100644 index 0000000..d083f4c --- /dev/null +++ b/templates/admin/users.html @@ -0,0 +1,34 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Пользователи{% endblock %} +{% block content %} +

Пользователи

+
+

Загрузка списка пользователей...

+ + + + + +
IDEmailИмяРольАктивенTelegramДействия
+
+ +{% endblock %} diff --git a/templates/auth/login.html b/templates/auth/login.html new file mode 100644 index 0000000..7b79f21 --- /dev/null +++ b/templates/auth/login.html @@ -0,0 +1,66 @@ +{% extends "base.html" %} + +{% block title %}Вход — ASR Сервис{% endblock %} + +{% block content %} +
+
+

Вход

+
+
+ + +
+
+ + +
+ +
+

+ Нет аккаунта? Зарегистрироваться +

+
+
+ + +{% endblock %} diff --git a/templates/auth/register.html b/templates/auth/register.html new file mode 100644 index 0000000..f580a81 --- /dev/null +++ b/templates/auth/register.html @@ -0,0 +1,72 @@ +{% extends "base.html" %} + +{% block title %}Регистрация — ASR Сервис{% endblock %} + +{% block content %} +
+
+

Регистрация

+
+
+ + +
+
+ + +
+
+ + +
+ +
+

+ Уже есть аккаунт? Войти +

+
+
+ + +{% endblock %} diff --git a/templates/base.html b/templates/base.html new file mode 100644 index 0000000..2cd2a61 --- /dev/null +++ b/templates/base.html @@ -0,0 +1,78 @@ + + + + + + + {% block title %}ASR Сервис{% endblock %} + + {% block head %}{% endblock %} + + + + + + +
+ {% block content %}{% endblock %} +
+ + + {% block scripts %}{% endblock %} + + diff --git a/templates/tg/index.html b/templates/tg/index.html new file mode 100644 index 0000000..a31947b --- /dev/null +++ b/templates/tg/index.html @@ -0,0 +1,41 @@ + + + + + + ASR Telegram + + + + +

ASR Сервис

+

Авторизация...

+ + + diff --git a/templates/user/api_keys.html b/templates/user/api_keys.html new file mode 100644 index 0000000..d6fc9b4 --- /dev/null +++ b/templates/user/api_keys.html @@ -0,0 +1,155 @@ +{% extends "user/base_user.html" %} + +{% block title %}API-ключи — ASR Сервис{% endblock %} + +{% block user_content %} +
+

API-ключи

+ +
+
+
+

Управление ключами

+

Ключи для программного доступа к API

+
+ +
+
+ +
+
+
+
+
+
+ + + +
+

Новый API-ключ

+
+ + +
+
+ + +
+
+
+ + + +
+

⚠️ Скопируйте ключ сейчас

+

Он больше не будет показан.

+
+ +
+
+ + +
+
+
+ + +{% endblock %} diff --git a/templates/user/asr.html b/templates/user/asr.html new file mode 100644 index 0000000..fd8a50d --- /dev/null +++ b/templates/user/asr.html @@ -0,0 +1,145 @@ +{% extends "user/base_user.html" %} + +{% block title %}ASR — Распознавание{% endblock %} + +{% block user_content %} +
+

Распознавание речи

+ {% if not user %} +
+
+ 👋 +
+
Гостевой режим
+
Вы можете распознавать аудио без регистрации (лимит 10 запросов/сутки). Зарегистрируйтесь, чтобы сохранять историю и увеличить лимиты.
+
+
+
+ {% endif %} + + +
+ + + +
+ + +
+
+
+ + +
+
+ + + + +
+ + +
+
+ + + + + + + + +
+

Последние сессии

+
+
+
+ + +{% endblock %} diff --git a/templates/user/base_user.html b/templates/user/base_user.html new file mode 100644 index 0000000..1f2c4e6 --- /dev/null +++ b/templates/user/base_user.html @@ -0,0 +1,44 @@ +{% extends "base.html" %} + +{% block content %} +
+ + + + + + +
+ {% block user_content %}{% endblock %} +
+
+ + +{% endblock %} diff --git a/templates/user/dashboard.html b/templates/user/dashboard.html new file mode 100644 index 0000000..300fa1a --- /dev/null +++ b/templates/user/dashboard.html @@ -0,0 +1,132 @@ +{% extends "user/base_user.html" %} + +{% block title %}Dashboard — ASR Сервис{% endblock %} + +{% block user_content %} +
+

Dashboard

+ +
+ +
+
+
+ + +
+
+
+ + +
+
+
+ + + +
+
+ + +{% endblock %} diff --git a/templates/user/history.html b/templates/user/history.html new file mode 100644 index 0000000..ca127fb --- /dev/null +++ b/templates/user/history.html @@ -0,0 +1,156 @@ +{% extends "user/base_user.html" %} + +{% block title %}История сессий — ASR Сервис{% endblock %} + +{% block user_content %} +
+

История сессий

+ + +
+
+
+ + +
+
+ + +
+ + +
+
+ + +
+
+
+
+ +
+
+ + +{% endblock %} diff --git a/templates/user/history_detail.html b/templates/user/history_detail.html new file mode 100644 index 0000000..7f542fe --- /dev/null +++ b/templates/user/history_detail.html @@ -0,0 +1,104 @@ +{% extends "user/base_user.html" %} + +{% block title %}Детали сессии — ASR Сервис{% endblock %} + +{% block user_content %} +
+

Детали сессии

+ +
+
+
+ + +
+ + +{% endblock %} diff --git a/templates/user/profile.html b/templates/user/profile.html new file mode 100644 index 0000000..f30af3f --- /dev/null +++ b/templates/user/profile.html @@ -0,0 +1,131 @@ +{% extends "user/base_user.html" %} + +{% block title %}Профиль — ASR Сервис{% endblock %} + +{% block user_content %} +
+

Профиль

+ +
+

Личные данные

+
+
+ + +
+
+ + +
+
+ + +
+ +
+
+ +
+

Смена пароля

+
+
+ + +
+
+ + +
+
+ + +
+ +
+
+ +
+

Опасная зона

+

Удаление аккаунта необратимо. Все данные будут сохранены, но доступ будет закрыт.

+ +
+
+ + +{% endblock %} diff --git a/templates/user/subscription.html b/templates/user/subscription.html new file mode 100644 index 0000000..cc3d1c1 --- /dev/null +++ b/templates/user/subscription.html @@ -0,0 +1,91 @@ +{% extends "user/base_user.html" %} + +{% block title %}Подписка — ASR Сервис{% endblock %} + +{% block user_content %} +
+

Управление подпиской

+ + +
+
+
+ + +

Доступные тарифы

+
+
+ Выбор и оплата тарифов — в разработке.
+ Перейти к распознаванию +
+
+ + +

История платежей

+
+
Нет данных
+
+
+ + +{% endblock %} diff --git a/to_do_front.txt b/to_do_front.txt new file mode 100644 index 0000000..43b5569 --- /dev/null +++ b/to_do_front.txt @@ -0,0 +1,631 @@ +# План разработки фронтенда и бэкенда ASR-сервиса + +--- + +## Этап -1. Подготовительный (без написания кода) + +- [x] Изучить и отформатировать текущий файл так, чтобы его было удобно читать. + +--- + +## Этап 0. Подготовительный (без написания кода) + +- [x] Проанализировать текущую кодовую базу и архитектуру. + +**Результат анализа:** + +### Архитектура FastAPI +- **Точка входа:** `main.py` (middleware `RequestIDMiddleware`, `DeprecationHeaderMiddleware`). +- **Middleware (по тестам):** CORS, GZip, ProxyHeaders, TrustedHost, RequestID (`tests/test_*middleware*.py`). +- **Конфигурация:** `config.py` — `Settings` на `pydantic-settings`, переменные окружения + `.env`. + +### Существующие endpoints +| Метод | Путь | Назначение | +|-------|------|------------| +| GET | `/` | Корневой (legacy) | +| GET | `/demo` | HTML-страница демо (`api/legacy/demo_page.py`) | +| POST | `/api/v1/asr/url` | Распознавание по URL (`api/legacy/post_by_url.py`) | +| POST | `/api/v1/asr/file` | Распознавание файла (POST) | +| WS | `/api/v1/asr/ws` | Потоковое распознавание (WebSocket) (`api/v1/endpoints/asr_ws.py`) | + +### Модели данных (Pydantic) +- **Домен:** `User`, `UserAccount`, `QuotaInfo`, `UserProfileResponse` (`models/domain/user.py`); `Subscription`, `Transaction`, `RecurringPayment` (`models/domain/billing.py`); `ASRUsageLog`, `AdminActionLog` (`models/domain/audit.py`). +- **API:** `BaseResponse`, `V1BaseResponse`, `ASRData`, `SyncASRRequest`, `PostFileRequest` и др. (`models/fast_api_models.py`). +- **WebSocket:** `WSConfigMessage`, `WSAudioMessage`, `WSPingMessage`, `WSStatusResponse` и др. (`models/ws_models.py`). Уже есть `parse_ws_message` и `wrap_binary_audio`. +- **ORM / БД:** В видимых файлах **нет** SQLAlchemy/Tortoise моделей или Alembic-миграций. Данные сессий хранятся в `InMemoryStateStore` / `RedisStateStore` (`core/state_store.py`). Для пользователей и подписок требуется выбрать и подключить ORM. + +### Авторизация +- **Заготовки:** `core/security.py` — хеширование `bcrypt`, создание/декодирование JWT (`access`/`refresh`), `TokenPayload`. Готов к интеграции. +- **Статус:** Endpoints ASR **не защищены** авторизацией. Зависимости `get_recognizer`, `get_punctuator`, `get_diarizer` используют `HTTPConnection`, но не проверяют JWT. + +### Платёжная система +- **Заготовки:** Pydantic-модели `Subscription`, `Transaction`, `RecurringPayment` есть, но таблиц в БД и API-эндпоинтов нет. + +### WebSocket / Мониторинг +- **Реализовано:** полный цикл в `services/ws_handler.py`, `ws_manager.py`, `ws_session.py`, `ws_metrics.py`. +- **Клиент:** `static/monitor.html` подключается к `/api/v1/asr/ws`, отправляет `status_request`, получает `status_response`. +- **Heartbeat:** добавлен в `monitor.html` (ping каждые 30 сек) для предотвращения `idle_timeout`. + +### План интеграции (уточнения) +- **ORM & БД:** SQLAlchemy 2.0 + Alembic + `asyncpg`. Создаём с нуля таблицы для пользователей, подписок, транзакций, ASR-сессий, системных логов. +- **Авторизация:** Access токен хранится в памяти клиента (memory), Refresh токен — в `httpOnly Secure SameSite=Strict` cookie. WebSocket-аутентификация — первый фрейм после connect с передачей токена. +- **Фронтенд:** Jinja2 + чистый HTML+JS (серверный рендеринг шаблонов). Существующие `templates/index.html` и `static/monitor.html` адаптируются под новую структуру. +- **Платежи:** Интеграция с ЮKassa. +- **Rate limiting:** По количеству запросов в минуту, с учётом тарифа пользователя. +- **API-доступ:** Поддержка API keys для программного доступа. +- Интегрировать `core/security.py` в FastAPI-зависимости (`Depends`) для защиты endpoint'ов. + +### Telegram Web App (дополнительный канал) +- **Контекст:** приложение будет доступно пользователям Telegram через Web App (встроенный WebView). +- **Авторизация:** Telegram передаёт `initData` (query string) с подписью HMAC-SHA256, проверяемой через `TELEGRAM_BOT_TOKEN`. Пользователь идентифицируется по `telegram_id`. +- **UI:** отдельный набор шаблонов/JS под мобильный viewport, использование `Telegram.WebApp` SDK (theme params, MainButton, BackButton, viewport stable height). +- **Платежи:** возможность использования Telegram Stars или `openInvoice` внутри Web App как альтернатива/дополнение к ЮKassa. +- **Бот:** требуется минимальный Telegram-бот для генерации ссылки на Web App и приёма обратных вызовов (menu button, inline keyboard). + +--- + +## Этап 1. Бэкенд: Расширение моделей данных [x] + +**Задача:** Создать с нуля SQLAlchemy 2.0 модели для поддержки авторизации, платежей, админки и API-доступа. + +### Инфраструктура БД +- Установить зависимости: `sqlalchemy[asyncio]`, `alembic`, `asyncpg`. +- Создать `db/base.py` — `DeclarativeBase` с общими `id` (UUID), `created_at`, `updated_at`. +- Создать `db/session.py` — `async_sessionmaker` + функция `get_db_session()` для Depends. +- Создать `db/enums.py` — SQLAlchemy-Enum для ролей, статусов, типов сессий. +- Инициализировать Alembic (`alembic init alembic`). +- Настроить `alembic.ini` и `alembic/env.py` для работы с `config.py:Settings.DATABASE_URL` (async). +- Сгенерировать первую миграцию (`alembic revision --autogenerate -m "init"`) и накатить (`alembic upgrade head`). + +### Модель пользователя (`User`) +- **Поля:** `id` (UUID/PK), `email` (unique, indexed), `hashed_password`, `full_name`, `phone` (optional), `role` (`user`/`admin`/`superadmin`), `is_active`, `email_verified`, `created_at`, `updated_at`, `last_login_at`. +- **Индексы:** составной по `role` + `is_active`. + +### Модель тарифного плана (`Plan`) — справочник +- **Поля:** `id`, `code` (unique), `name`, `description`, `max_requests_per_minute`, `max_audio_duration_sec`, `price_per_month`, `is_active`. +- **Назначение:** хранение лимитов для rate limiting и цен для ЮKassa. + +### Модель подписки (`Subscription`) +- **Поля:** `id`, `user_id` (FK → `users.id`), `plan_id` (FK → `plans.id`), `status` (`active`/`expired`/`cancelled`), `started_at`, `expires_at`, `auto_renew`, `yookassa_payment_id`. +- **Связи:** `user`, `plan`. + +### Модель транзакции (`Transaction`) +- **Поля:** `id`, `user_id` (FK), `subscription_id` (FK, nullable), `amount`, `currency` (`RUB`), `status` (`pending`/`succeeded`/`cancelled`), `payment_provider` (`yookassa`), `external_payment_id`, `created_at`, `metadata` (JSON). +- **Индекс:** по `external_payment_id` для обработки webhooks. + +### Модель сессии ASR (`ASRSession`) +- **Поля:** `id`, `user_id` (FK, nullable для совместимости), `session_type` (`url`/`file`/`websocket`), `status` (`processing`/`completed`/`failed`), `audio_duration_sec`, `processing_duration_sec`, `cost` (numeric, nullable), `result_json` (JSONB, nullable), `created_at`, `completed_at`, `error_message`, `request_ip`, `user_agent`. +- **Индексы:** по `user_id`, `created_at` (desc), `status`. + +### Модель API-ключа (`ApiKey`) +- **Поля:** `id`, `user_id` (FK), `name`, `key_hash` (bcrypt/sha256), `permissions` (JSON), `is_active`, `rate_limit_override` (int, nullable), `created_at`, `last_used_at`. +- **Индексы:** по `user_id`, `key_hash` (unique). + +### Модель системного лога (`SystemLog`) +- **Поля:** `id`, `level` (`info`/`warning`/`error`), `component`, `message`, `metadata` (JSON), `created_at`. +- **Индекс:** по `created_at` (desc), `level`, `component`. + +### Модель аудита админов (`AdminAuditLog`) +- **Поля:** `id`, `admin_id` (FK → users.id), `action`, `target_type`, `target_id`, `details` (JSON), `created_at`. +- **Индекс:** по `admin_id`, `created_at` (desc). + +### Дополнение модели `User` (Telegram) +- **Поля:** `telegram_id` (bigint, unique, nullable), `telegram_username` (string, nullable), `telegram_first_name` (string, nullable), `telegram_last_name` (string, nullable), `telegram_photo_url` (text, nullable), `telegram_auth_date` (datetime, nullable). +- **Индекс:** по `telegram_id`. + +### Модель конфигурации Telegram-бота (`TelegramBotConfig`) — опционально +- **Поля:** `id`, `bot_token_hash` (для хранения токена бота, если нужно в БД), `webapp_url`, `is_active`, `created_at`. +- **Назначение:** хранение настроек бота для админ-панели (можно заменить env-переменными). + +**Результат этапа 1:** созданы `db/base.py`, `db/session.py`, `db/enums.py`, `db/models.py`, `alembic.ini`, `alembic/env.py`, `alembic/script.py.mako`, `alembic/README`. SQLAlchemy 2.0 + aiosqlite + Alembic настроены, модели описаны (включая Telegram-поля в `User` и `TelegramBotConfig`), первая миграция выполнена. Готово к генерации миграции для Telegram-полей. + +--- + +## Этап 2. Бэкенд: Система авторизации, ролей и API-ключей [x] + +**Задача:** Реализовать двухфакторную схему JWT (access/refresh) + API Key auth для программного доступа. Использовать существующий `core/security.py` как основу. + +### JWT Аутентификация (веб) +- Реализовать в `api/v1/endpoints/auth.py`: + | Метод | Endpoint | Описание | + |-------|----------|----------| + | POST | `/api/v1/auth/register` | Регистрация (`email`, `password`, `full_name`). Access в JSON, refresh в `httpOnly` cookie. | + | POST | `/api/v1/auth/login` | Вход. Access token в теле ответа, refresh token — в `httpOnly Secure SameSite=Strict` cookie. | + | POST | `/api/v1/auth/refresh` | Обмен валидного refresh cookie на новую пару access/refresh. Проверка blacklist. | + | POST | `/api/v1/auth/logout` | Инвалидация refresh токена (чёрный список в `RedisStateStore`/`InMemoryStateStore`) + очистка cookie. | + | POST | `/api/v1/auth/change-password` | Смена пароля (требует старый пароль + access token). | + | GET | `/api/v1/auth/me` | Текущий пользователь по access token. | + +### WebSocket-аутентификация +- Первый фрейм после `connect` обязан содержать сообщение типа `auth`: `{"type":"auth","access_token":"..."}`. +- Использовать `parse_ws_message` из `models/ws_models.py` для разбора. +- Сервер валидирует токен через `core/security.py` и ассоциирует `AudioSession` / `ConnectionMeta` с `user_id`. +- При невалидном токене — немедленный `close` с кодом `1008` (Policy Violation). +- Для публичных WS (если останутся) — `user_id = null`. + +### API Key аутентификация (программный доступ) +- Заголовок `X-API-Key: `. +- Dependency `require_api_key` — проверка хеша ключа в БД (`ApiKey`), проверка `is_active`, обновление `last_used_at`. +- API keys действуют только на REST endpoints ASR (не на WS). +- Возможность ограничить permissions ключа (например, только `asr:read`). + +### Авторизация (RBAC) +- Зависимости FastAPI в `core/deps.py` (новый файл): `require_auth` (JWT), `require_admin`, `require_superadmin`, `require_api_key_or_auth`. +- Middleware для извлечения `current_user` из JWT (cookie/header) или API Key (header). +- ASR HTTP endpoints: `require_auth` или `require_api_key_or_auth`. +- ASR WS: auth через первый фрейм (см. выше). +- Админ endpoints: `require_admin`. + +### Прочее +- **Хеширование паролей:** `bcrypt` (через `passlib`), переиспользовать `core/security.py`. +- **Хеширование API keys:** `bcrypt` или `sha256` + salt (ключ показывается пользователю только один раз при создании). +- **Валидация:** Pydantic v2 схемы для всех auth-запросов (`models/fast_api_models.py` или отдельный `models/auth.py`). +- **Logout / инвалидация:** refresh токены хранятся в `RedisStateStore`/`InMemoryStateStore` с TTL, возможность отзыва. + +### Telegram Web App аутентификация +- **Механизм:** при открытии Web App клиент получает `initData` из `Telegram.WebApp.initData`. Отправляет его на `POST /api/v1/auth/telegram`. +- **Валидация:** сервер проверяет HMAC-SHA256 подписи `initData` с использованием `TELEGRAM_BOT_TOKEN` (SHA256 хеш токена — секрет). +- **Поведение:** если `telegram_id` не найден — автоматическая регистрация (`User` с `role=user`, без пароля). Если найден — выдача access/refresh токенов (так же, как при обычном логине). +- **Связь с существующей системой:** Telegram-авторизация не отменяет email/password auth, а дополняет её. Пользователь может иметь оба метода. + +**Результат этапа 2:** созданы `models/auth.py`, `core/deps.py`, `services/auth_service.py`, `api/v1/endpoints/auth.py`. Реализованы JWT (access/refresh в cookie), API Key auth, RBAC-зависимости, Telegram Web App auth с проверкой HMAC, blacklist refresh-токенов в памяти. + +--- + +## Этап 3. Бэкенд: API для пользовательского кабинета [x] + +**Задача:** Создать endpoints для работы обычных пользователей в `api/v1/endpoints/user.py`. + +### Профиль пользователя +- `GET /api/v1/user/profile` — получить профиль (из БД). +- `PUT /api/v1/user/profile` — обновить профиль (имя, телефон). +- `POST /api/v1/user/change-password` — смена пароля (отдельно от профиля). +- `DELETE /api/v1/user/profile` — удалить аккаунт (soft delete, `is_active = false`). + +### Квота и подписка +- `GET /api/v1/user/quota` — остаток запросов/секунд в текущем периоде (расчёт на основе `Plan` и использования). +- `GET /api/v1/user/subscription` — текущая подписка и детали плана. +- `POST /api/v1/user/subscription/upgrade` — запрос на смену тарифа (заглушка под платёжный провайдер). +- `POST /api/v1/user/subscription/cancel` — отмена авто-продления. + +### История использования ASR +- `GET /api/v1/user/sessions` — список сессий распознавания (пагинация, фильтры по дате, статусу). +- `GET /api/v1/user/sessions/{session_id}` — детали конкретной сессии + `result_json`. +- `GET /api/v1/user/stats` — агрегированная статистика (всего часов обработано, за месяц, количество сессий). + +### API-ключи +- `GET /api/v1/user/api-keys` — список ключей (без plain key, только имена и даты). +- `POST /api/v1/user/api-keys` — создать ключ (вернуть plain key один раз). +- `DELETE /api/v1/user/api-keys/{key_id}` — отозвать ключ (`is_active = false`). + +### Telegram Web App endpoints +- `GET /api/v1/user/telegram/link` — проверка, привязан ли Telegram-аккаунт (для email-пользователей). +- `POST /api/v1/user/telegram/unlink` — отвязать Telegram (очистка `telegram_id` и т.д., требует подтверждение паролем/email). +- `GET /tg` — серверно-рендеримая точка входа для Telegram Web App (`templates/tg/index.html`). Должна принимать `initData` и сразу обменивать его на JWT через JS. + +**Результат этапа 3:** созданы `models/user.py`, `api/v1/endpoints/user.py`, `api/v1/endpoints/tg.py`, `templates/tg/index.html`. Реализованы профиль, квота, подписка, история ASR, API-ключи, статистика, Telegram link/unlink, точка входа `/tg` для Web App. + +--- + +## Этап 4. Бэкенд: API и Jinja-шаблоны админ-панели [x] + +**Задача:** Создать endpoints и страницы для мониторинга и управления системой. + +> **Стек фронта:** Jinja2 + чистый HTML/JS. Админ-страницы под `/admin/*`, API под `/api/v1/admin/*`. + +### Мониторинг системы +- `GET /admin/dashboard` (HTML) / `GET /api/v1/admin/metrics` (JSON) — текущие метрики из `SystemMetricsCollector`. +- `GET /api/v1/admin/metrics/history` — история метрик за период (для графиков Chart.js). +- `WS /api/v1/admin/ws` — real-time поток метрик. **Auth:** первый фрейм `{"type":"auth","access_token":"..."}`. Проверка `require_admin`. Использовать `MessageRouter` из `services/ws_handler.py`. + +### Управление пользователями +- `GET /admin/users` (HTML) / `GET /api/v1/admin/users` (JSON) — пагинация, поиск, фильтры. +- `GET /api/v1/admin/users/{user_id}` — детали. +- `PUT /api/v1/admin/users/{user_id}` — редактирование роли, статуса, блокировки. +- `DELETE /api/v1/admin/users/{user_id}` — soft delete / блокировка. +- `GET /api/v1/admin/users/{user_id}/sessions` — история ASR-сессий пользователя. +- `POST /api/v1/admin/users/{user_id}/impersonate` — получить access token от имени пользователя (superadmin only). + +### Управление тарифами (справочник) +- `GET /admin/tariffs` (HTML) / `GET /api/v1/admin/tariffs` (JSON) — список тарифных планов. +- `POST /api/v1/admin/tariffs` — создать тариф. +- `PUT /api/v1/admin/tariffs/{plan_id}` — изменить лимиты/цену. +- `DELETE /api/v1/admin/tariffs/{plan_id}` — деактивировать тариф (`is_active = false`). + +### Управление подписками и платежами +- `GET /admin/subscriptions` (HTML) / `GET /api/v1/admin/subscriptions` (JSON). +- `GET /admin/transactions` (HTML) / `GET /api/v1/admin/transactions` (JSON). +- `POST /api/v1/admin/subscriptions/{sub_id}/extend` — ручное продление (без оплаты). +- `POST /api/v1/admin/subscriptions/{sub_id}/cancel` — ручная отмена. + +### Логи и аудит +- `GET /admin/logs` (HTML) / `GET /api/v1/admin/logs` (JSON) — системные логи. +- `GET /api/v1/admin/audit` — аудит действий админов. + +### Управление ASR-очередью и сессиями +- `GET /api/v1/admin/queue` — текущая очередь задач (из `ws_metrics` / state store). +- `POST /api/v1/admin/queue/{task_id}/cancel` — отменить задачу. +- `POST /api/v1/admin/users/{user_id}/sessions/{session_id}/disconnect` — принудительно закрыть WS сессию (через `ConnectionManager.disconnect`). +- `POST /api/v1/admin/maintenance` — включить/выключить режим обслуживания (запретить новые ASR-сессии). + +### Управление API-ключами пользователей +- `GET /api/v1/admin/api-keys` — все ключи (с фильтрами по пользователю). +- `DELETE /api/v1/admin/api-keys/{key_id}` — отозвать ключ админом. + +**Результат этапа 4:** созданы `models/admin.py`, `services/admin_service.py`, `api/v1/endpoints/admin.py`, `routes/admin.py`, шаблоны `templates/admin/base_admin.html`, `dashboard.html`, `login.html`, `users.html`, `sessions.html`, `subscriptions.html`, `transactions.html`, `tariffs.html`, `api_keys.html`, `logs.html`, `settings.html`, `telegram.html`. Реализованы API и HTML-роуты для метрик, пользователей, тарифов, подписок, транзакций, логов, аудита, API-ключей, управление очередью, maintenance mode, impersonate, Telegram-конфигурация. Все админ-страницы защищены `require_admin`. + +**Результат этапа 4 (завершение):** в `main.py` подключены все новые роутеры (`auth`, `user`, `admin`, `tg`, `routes.admin`). Этап полностью выполнен. + +--- + +## Этап 5. Бэкенд: Интеграция ASR, мониторинга, rate limiting и пользователей [x] + +**Задача:** Связать существующий ASR-конвейер с БД, авторизацией и лимитами. + +### Привязка ASR-сессий к пользователям +- HTTP endpoints (`/api/v1/asr/url`, `/api/v1/asr/file`): извлекать `current_user` из JWT/API Key и записывать `user_id` в `ASRSession`. +- Интегрировать создание `ASRSession` в `Recognizer/engine/file_recognition.py` и `services/asr_pipeline.py`. +- WebSocket `/api/v1/asr/ws`: после auth-фрейма записывать `user_id` в `AudioSession` (`services/ws_session.py`) и `ConnectionMeta` (`services/ws_manager.py`). Если auth не пришёл в течение 5 сек — `close(1008)`. +- Сохранять `result_json` в `ASRSession` после завершения обработки (в `do_sensitizing`/`process_file`). + +### Rate limiting +- **Стратегия:** по количеству запросов в минуту на основе тарифа (`Plan.max_requests_per_minute`). +- **Реализация:** `RateLimitMiddleware` в `main.py` или dependency `check_rate_limit`. Хранилище счётчиков — `RedisStateStore` (предпочтительно) или `InMemoryStateStore` (для single-node). +- **Ключ лимита:** `rate_limit:{user_id}` или `rate_limit:{api_key_hash}`. +- **Ответ при превышении:** `429 Too Many Requests` с заголовком `Retry-After`. +- **Исключения:** админы и superadmin не подлежат rate limiting. + +### Сбор метрик в БД и Prometheus +- Фоновая задача (`asyncio` background task или APScheduler) — писать срезы метрик из `SystemMetricsCollector` в таблицу `SystemLog` каждые 30 сек. +- Опционально: интеграция с `prometheus-client` + endpoint `/metrics` для Grafana. + +### WebSocket для админов +- Endpoint `/api/v1/admin/ws`. +- **Auth:** первый фрейм `{"type":"auth","access_token":"..."}` → проверка `require_admin`. +- При успешной авторизации — отправка текущего статуса, затем push каждые 5 сек (или по событию) из `SystemMetricsCollector`. +- **Формат сообщений:** + ```json + { "type": "metrics", "data": {...} } + { "type": "alert", "data": {...} } + ``` +- Отображение в `admin/dashboard` (виджет System Health). + +**Результат этапа 5:** созданы `core/middleware.py` (RateLimitMiddleware с maintenance mode, 429 + Retry-After), `services/metrics_reporter.py` (фоновая запись метрик в `SystemLog` каждые 30 сек), `api/v1/endpoints/admin_ws.py` (WebSocket `/api/v1/admin/ws` с auth фреймом и push метрик каждые 5 сек). В `api/v1/endpoints/asr_ws.py` добавлена обязательная WS-аутентификация (auth фрейм, таймаут 5 сек, привязка `user_id` к `AudioSession`) и создание/обновление `ASRSession` в БД с `result_json`. Все компоненты подключены в `main.py`. + +**Для полной привязки ASR-сессий к `user_id` в HTTP endpoints** необходимо добавить в чат и отредактировать: `services/asr_pipeline.py`, `Recognizer/engine/file_recognition.py`. + +--- + +## Этап 6. Фронтенд: Пользовательский кабинет (Jinja2 + JS) + +**Задача:** Создать серверно-рендеримые страницы для конечных пользователей с единой дизайн-системой, авторизацией через JWT и поддержкой Telegram Web App. + +### Шаг 6.0. Инфраструктура (фундамент) [x] +**Цель:** Создать JS-ядро, CSS-переменные и базовые layout'ы, которые будут использоваться во всех страницах. + +| Файл | Действие | +|------|----------| +| `static/css/design-system.css` | Создать. CSS-переменные: `--color-bg: #0f172a`, `--color-surface: #1e293b`, `--color-primary: #3b82f6`, `--color-text: #e2e8f0`, `--color-muted: #94a3b8`, `--color-success: #22c55e`, `--color-warning: #f59e0b`, `--color-danger: #ef4444`. Тёмная тема по умолчанию. Утилиты: `.card`, `.btn`, `.btn-primary`, `.btn-danger`, `.btn-ghost`, `.form-group`, `.input`, `.badge`, `.skeleton`, `.toast-container`, `.spinner`. | +| `static/js/auth.js` | Создать. Модуль с замыканием: `let _token = null;`. Функции: `setAccessToken(t)`, `getAccessToken()`, `clearAuth()`, `apiFetch(url, opts)` — автоподстановка `Authorization: Bearer`, автоматический `/auth/refresh` при 401, редирект на `/login` при неудаче refresh, обработка 403. Функция `initAuth()` — вызывается на каждой странице, проверяет наличие токена. | +| `static/js/ui.js` | Создать. `toast(msg, type='info', duration=5000)` — fixed-контейнер в правом верхнем углу, auto-dismiss с прогресс-баром. `confirmDialog(title, text, onConfirm)` — нативный ``, не `window.confirm()`. `setLoading(element, isLoading)` — добавляет/убирает `.spinner` и `disabled`. `formatDate(iso)`, `formatBytes(bytes)`, `formatDuration(sec)`. | +| `templates/base.html` | Создать. Блоки: `title`, `head`, `nav`, `content`, `scripts`. Подключение `design-system.css`, `auth.js`, `ui.js`. CSRF-токен в ``. Навигация: логотип, ссылки (если авторизован: Dashboard, ASR, History, Subscription, Profile, Logout), иначе (Login, Register). | +| `templates/user/base_user.html` | Создать. Наследует `base.html`. Добавляет боковое меню пользователя (Desktop: sidebar 240px, Mobile: bottom tab-bar или hamburger). Подсветка активного пункта через `request.url.path`. Блок `user_content`. | +| `main.py` (или роуты) | Добавить HTML-роуты: `GET /login`, `GET /register`, `GET /dashboard`, `GET /asr`, `GET /history`, `GET /history/{id}`, `GET /subscription`, `GET /profile`, `GET /api-keys`. Все защищены `require_auth`, кроме `/login` и `/register`. | + +**Контракт apiFetch:** +```javascript +// При 401 делает POST /api/v1/auth/refresh, затем retry исходного запроса +// При неудаче refresh — clearAuth() + window.location.href = '/login' +// Всегда добавляет header Authorization: Bearer +// Возвращает Promise +``` + +**Результат шага 6.0:** созданы `static/css/design-system.css`, `static/js/auth.js`, `static/js/ui.js`, `templates/base.html`, `templates/user/base_user.html`, `routes/user.py`. В `main.py` подключены HTML-роуты пользовательского кабинета. + +### Шаг 6.1. Аутентификация (Login / Register) [x] +**Цель:** Страницы входа и регистрации, которые получают access token и сохраняют его в auth.js. + +| Файл | Действие | +|------|----------| +| `templates/auth/login.html` | Наследует `base.html` (без user-меню). Центрированная карточка. Поля: email, password. Кнопка «Войти». Ссылка на регистрацию. JS: `handleLogin` — `fetch('/api/v1/auth/login', {credentials:'include'})`, при успехе сохраняет `access_token` через `setAccessToken()`, редирект на `/dashboard`. Ошибки — `toast()` под формой. | +| `templates/auth/register.html` | Аналогично. Поля: email, password, full_name. После регистрации — автоматический вход или редирект на `/login` с `toast('Регистрация успешна')`. | + +**Важно:** Формы отправляются через обычный `fetch` (не HTMX), т.к. нужно программно сохранить токен из JSON-ответа. + +**Результат шага 6.1:** созданы `templates/auth/login.html` и `templates/auth/register.html`. Обе страницы наследуют `base.html`, используют `auth.js` (setAccessToken) и `ui.js` (toast, setLoading). При успешном входе/регистрации — редирект на `/dashboard`. + +### Шаг 6.2. Dashboard (главная панель) [x] +**Цель:** Обзорная страница с квотой, статистикой и быстрыми действиями. + +| Файл | Действие | +|------|----------| +| `templates/user/dashboard.html` | Наследует `user/base_user.html`. Сетка виджетов (grid 1-4 колонки). Виджеты: (1) Текущий тариф (badge), остаток квоты; (2) Статистика за месяц (количество сессий, минут аудио); (3) Последние 5 сессий (мини-таблица, загружается через `apiFetch`); (4) Быстрые ссылки (кнопки: «Новое распознавание», «История», «API-ключи»). | + +**UX:** +- Загрузка: skeleton-заглушки в виджетах, не спиннер на весь экран. +- Ошибка: toast + retry-кнопка на виджете. +- Пустое состояние: «Нет сессий. Начните с ASR» с иконкой. + +**Результат шага 6.2:** создан `templates/user/dashboard.html`. Виджеты загружают данные через `Auth.apiFetch` с `/api/v1/user/quota`, `/api/v1/user/stats`, `/api/v1/user/sessions?limit=5`. Skeleton-заглушки заменяются на контент при загрузке. При ошибке — toast + кнопка повтора. Пустое состояние — ссылка на `/asr`. + +### Шаг 6.3. ASR-интерфейс (самый сложный) [x] +**Цель:** Перенести весь функционал `templates/index.html` (legacy demo) в авторизованный интерфейс с современным UX. + +| Файл | Действие | +|------|----------| +| `static/js/asr_client.js` | Создать. Класс `ASRClient`. Методы: `connect(authToken, config)` — открывает WS `/api/v1/asr/ws`, отправляет auth-фрейм, затем config. `sendAudioChunk(base64Chunk, seqNum)` — отправляет `{"type":"audio_chunk",...}`. `sendEOS()` — конец потока. `onPartial(cb)`, `onFinal(cb)`, `onError(cb)`, `disconnect()`. Reconnect с exponential backoff (1с, 2с, 4с, max 30с). | +| `templates/user/asr.html` | Наследует `user/base_user.html`. Три вкладки/секции: (1) «По ссылке» — input URL + чекбоксы (keep_raw, do_echo_clearing, do_dialogue, do_punctuation) + кнопка «Отправить». Результат — карточка с текстом. (2) «Загрузка файла (POST)» — drag-and-drop зона + чекбоксы (do_diarization, diar_vad_sensity) + кнопка. Прогресс загрузки. (3) «WebSocket (WAV)» — input file (только .wav) + чекбоксы (ws_do_dialogue, ws_do_punctuation, ws_use_base64) + кнопка «Отправить». Область результата с partial/final текстом. | + +**UX-детали:** +- **Drag-and-drop:** зона с пунктирной рамкой, подсветка при dragover, иконка загрузки. Файл можно бросить на любую из трёх секций (определяет тип обработки). +- **WebSocket:** индикатор статуса соединения (dot: красный/жёлтый/зелёный). Кнопка «Стоп» для прерывания. +- **Результат:** текстовая область с кнопкой «Копировать» и «Скачать .txt». Для WS — накопление partial-результатов с таймкодами. +- **История рядом:** под основным блоком — последние 5 сессий пользователя (загружается через `apiFetch`), клик переходит на `/history/{id}`. + +**Адаптация из `index.html`:** +- Скопировать логику `readWavHeader`, `arrayBufferToBase64`, `sendChannel`, `sendAllChannels` из `index.html` в `asr_client.js`. +- Сохранить все существующие чекбоксы и их логику. +- Добавить `channel_name` как в оригинале. + +**Результат шага 6.3:** созданы `static/js/asr_client.js` (класс ASRClient с reconnect, auth-фрейм, partial/final/error коллбеки) и `templates/user/asr.html` (три вкладки: URL, POST-файл с drag-and-drop, WS WAV с индикатором статуса и кнопкой Стоп). Все секции используют `Auth.apiFetch` и `UI.toast`. История сессий подгружается под блоком ASR. + +### Шаг 6.4. История сессий [x] +**Цель:** Таблица с пагинацией, фильтрами и деталями. + +| Файл | Действие | +|------|----------| +| `templates/user/history.html` | Наследует `user/base_user.html`. Фильтры: дата (с/по), тип (url/file/websocket), статус. Таблица: дата, тип, длительность, статус (badge), действия (детали, скачать JSON). Пагинация: «Загрузить ещё» (кнопка) или номера страниц. `apiFetch` + JS для сложной фильтрации. | +| `templates/user/history_detail.html` | Наследует `user/base_user.html`. Карточка с метаданными (дата, тип, длительность, IP). Блок результата: pretty-printed JSON в `
` с подсветкой синтаксиса (можно просто CSS) или текст распознавания. Кнопки: «Скачать JSON», «Скачать TXT», «Назад к истории». |
+
+**UX:**
+- **Статусы:** completed (зелёный), processing (жёлтый пульсирующий), failed (красный).
+- **Скелетон** при загрузке.
+- **Empty state:** иконка + «Нет сессий. Перейдите в ASR».
+
+**Результат шага 6.4:** созданы `templates/user/history.html` и `templates/user/history_detail.html`. Реализованы фильтры по типу/статусу, пагинация «Загрузить ещё», skeleton, empty state, детали сессии с pretty JSON и кнопками скачивания.
+
+### Шаг 6.5. Подписка, профиль, API-ключи [x]
+**Цель:** CRUD-страницы для управления аккаунтом.
+
+| Файл | Действие |
+|------|----------|
+| `templates/user/subscription.html` | Текущий тариф (карточка), дата окончания, auto-renew (toggle). Список доступных тарифов (карточки с ценой, кнопка «Выбрать» — заглушка до интеграции ЮKassa). История платежей (мини-таблица). |
+| `templates/user/profile.html` | Форма: full_name (редактируемое), email (read-only или с верификацией), смена пароля (old + new + confirm). Кнопка «Удалить аккаунт» — `confirmDialog` + soft-delete. |
+| `templates/user/api_keys.html` | Таблица ключей: имя, префикс (последние 4 символа), дата создания, last_used_at, статус. Кнопка «Создать ключ» — модальное окно с именем. При создании: показать ключ один раз в модалке с кнопкой «Копировать». Кнопка «Отозвать» — `confirmDialog`. |
+
+**UX:** [не валидировано - требуется регистрация]
+- **API-ключ при создании:** модалка с жёлтым предупреждением «Скопируйте сейчас, ключ больше не будет показан».
+- **Профиль:** inline-валидация (minlength для пароля).
+
+**Результат шага 6.5:** созданы `templates/user/subscription.html`, `templates/user/profile.html`, `templates/user/api_keys.html`. Все страницы наследуют `user/base_user.html`, используют `Auth.apiFetch` и `UI.toast`/`UI.confirmDialog`.
+
+### Шаг 6.6. Telegram Web App
+**Цель:** Адаптация интерфейса под TG Web App с нативными кнопками.
+
+| Файл | Действие |
+|------|----------|
+| `templates/tg/base_tg.html` | Создать. Минимальный layout: подключение `telegram-web-app.js`, `Telegram.WebApp.expand()`, viewport meta, цвета темы из `Telegram.WebApp.themeParams`. Нет бокового меню — только экраны. |
+| `templates/tg/index.html` | Изменить существующий. Экран входа: кнопка «Войти через Telegram» (отправка initData на `/api/v1/auth/telegram`), затем редирект на `/tg/asr`. |
+| `templates/tg/asr.html` | Создать. Упрощённый ASR: только загрузка файла (drag-and-drop не нужен, используем нативный input). `Telegram.WebApp.MainButton.setText('Распознать')` — запускает отправку. `MainButton.show()` при выборе файла. Результат — в блоке + кнопка «Поделиться». |
+| `templates/tg/history.html` | Создать. Список сессий (карточки, не таблица). `BackButton` для возврата. |
+| `templates/tg/profile.html` | Создать. Квота, тариф, кнопка «Закрыть» (`Telegram.WebApp.close()`). |
+| `static/js/tg_app.js` | Создать. Инициализация TG, обработка `themeChanged`, управление `MainButton`/`BackButton`. |
+
+---
+
+## Этап 7. Фронтенд: Админ-панель (Jinja2 + JS)
+
+**Задача:** Доработать серверно-рендеримый интерфейс администрирования с единой дизайн-системой, real-time мониторингом и полноценным CRUD.
+
+### Шаг 7.0. Инфраструктура админки
+**Цель:** Унифицировать UX админки с пользовательским фронтом.
+
+| Файл | Действие |
+|------|----------|
+| `static/js/admin_ui.js` | Создать. Расширяет `ui.js`: `adminToast`, `adminConfirm`, `renderBadge(status)`, `renderDate(iso)`. Функция `initAdmin()` — проверка роли admin при загрузке. |
+| `templates/admin/base_admin.html` | Изменить. Подключить `design-system.css` (для единых переменных), `admin_ui.js`. Добавить `
`. Добавить CSRF-meta. | + +### Шаг 7.1. Dashboard — real-time + история метрик +**Цель:** Заменить статический fetch на живой мониторинг. + +| Файл | Действие | +|------|----------| +| `templates/admin/dashboard.html` | Изменить. (1) Виджеты системных метрик подключаются к `/api/v1/admin/ws` (WebSocket с auth-фреймом). Данные обновляются в реальном времени. (2) Добавить Chart.js график «История нагрузки за 24ч»: endpoint `/api/v1/admin/metrics/history?range=24h` (заменить заглушку на реальную выборку из `SystemLog` за последние 24 часа, агрегируя по 5-минутным интервалам). | +| `static/js/admin_dashboard.js` | Создать. WS-клиент для `/api/v1/admin/ws`: подключение, auth-фрейм, heartbeat (ping каждые 30с), обработка `status_response`, обновление DOM-виджетов, reconnect. Алерты: если `gpu_utilization > 90` или `queue_depth > 50` — показать баннер «Высокая нагрузка». | + +**Бэкенд (если нужно):** +- Реализовать `/api/v1/admin/metrics/history` — выборка из `SystemLog` где `component = 'metrics'` и `created_at > now - range`, группировка по временным слотам. + +### Шаг 7.2. CRUD тарифов +**Цель:** Заменить readonly-таблицу на полноценное управление. + +| Файл | Действие | +|------|----------| +| `templates/admin/tariffs.html` | Изменить. Добавить кнопку «Создать тариф». Таблица: код, название, лимит запросов, лимит аудио, цена, статус. Действия: редактировать, деактивировать. | +| `templates/admin/components/tariff_modal.html` | Создать. `` с формой: code, name, max_requests_per_minute, max_audio_duration_sec, price_per_month, is_active (toggle). Отправка через `fetch` (POST/PUT). При успехе — закрытие модалки + обновление таблицы. | + +**UX:** +- При деактивации — `confirmDialog` («Деактивировать тариф "Базовый"? Существующие подписки останутся активными»). +- Валидация: code — уникальный, латиница + underscore. + +### Шаг 7.3. Детали пользователя +**Цель:** Страница детального просмотра и редактирования. + +| Файл | Действие | +|------|----------| +| `templates/admin/user_detail.html` | Создать. Роут `/admin/users/{id}`. Карточка пользователя: email, имя, роль (select: user/admin/superadmin), статус (toggle is_active), telegram, дата регистрации. Кнопки: «Сохранить», «Войти как пользователь» (impersonate, superadmin only), «Блокировать». Таблица сессий пользователя (последние 20, с пагинацией). Таблица подписок. | +| `templates/admin/users.html` | Изменить. Сделать строки кликабельными (переход на `/admin/users/{id}`). | + +**UX:** +- **Impersonate:** кнопка видна только superadmin, при клике — получение токена и открытие `/dashboard` в новой вкладке от имени пользователя. +- **Блокировка:** `confirmDialog` + soft-delete. + +### Шаг 7.4. Мониторинг сессий +**Цель:** Заменить заглушку `/admin/sessions` на реальные данные. + +| Файл | Действие | +|------|----------| +| `templates/admin/sessions.html` | Изменить. Таблица активных ASR-сессий: ID, пользователь (email), тип, статус, длительность аудио, начало. Кнопка «Отключить WS» (для websocket-сессий) — вызов `/api/v1/admin/users/{uid}/sessions/{sid}/disconnect`. Автообновление каждые 10 сек через `apiFetch`. | + +**Бэкенд (если нужно):** +- Реализовать хранение активных WS-сессий в `ConnectionMeta` с возможностью их перечисления и закрытия. + +### Шаг 7.5. UX/UI полировка всех админ-страниц +**Цель:** Единообразие, состояния, обработка ошибок. + +| Файл | Действие | +|------|----------| +| `templates/admin/users.html` | Добавить skeleton при загрузке. Поиск с debounce (300мс). Фильтры по роли/статусу (табы). Empty state. | +| `templates/admin/subscriptions.html` | Добавить фильтры по статусу и тарифу. Кнопки «Продлить»/«Отменить» с `confirmDialog`. | +| `templates/admin/transactions.html` | Добавить фильтр по статусу. | +| `templates/admin/logs.html` | Добавить фильтры по level (info/warning/error) и component. Live-обновление (опционально). | +| `templates/admin/api_keys.html` | Добавить поиск по user_id/email. | +| `templates/admin/settings.html` | Maintenance mode toggle — добавить состояние loading на кнопку, toast при успехе. | + +**Общие правила для всех админ-страниц:** +- Все кнопки, вызывающие `fetch`, имеют состояние loading (спиннер + disabled). +- Все опасные действия (delete, cancel, block) требуют `confirmDialog`. +- Все ошибки API отображаются через toast (красный). +- Все успешные действия — toast (зелёный). +- Таблицы: hover-эффект на строках, сортировка по клику на заголовок (если применимо). + +--- + +## Этап 8. Интеграция и тестирование + +**Задача:** Собрать всё вместе и проверить работоспособность. + +### Middleware и CORS +- Настроить CORS для разрешённых origin (фронтенд на том же домене — CORS минимален). +- Middleware логирования запросов (с `request_id`). +- `RateLimitMiddleware` в `main.py` с разными лимитами для `user`/`admin`/`api_key`. + +### Обработка ошибок +- Единый формат ошибок API: + ```json + { "detail": "...", "code": "..." } + ``` +- Расширить `core/exception_handlers.py` для новых исключений (RateLimitExceededException и т.д.). +- Отображение ошибок на Jinja-страницах через toast-уведомления (`static/js/ui.js`). + +### Тестирование +- **Unit:** auth endpoints, admin CRUD, tariff CRUD, API Key validation (`pytest`). +- **Интеграционные:** регистрация → логин → создание API ключа → распознавание через API key → просмотр истории. +- **WebSocket:** + - ASR WS: подключение → auth фрейм → config → audio → результат. + - Admin WS: подключение → auth фрейм → получение метрик. + - Heartbeat: проверка `ping`/`pong`. +- **Rate limiting:** превышение лимита → `429` с `Retry-After`. +- **Роли:** пользователь не может зайти в `/admin`, админ видит все данные, superadmin может имперсонировать. +- **Платежи:** эмуляция webhook ЮKassa (тестовый shopId). +- **Аудит:** проверка записей в `AdminAuditLog` после действий админа. + +### Документация API +- Обновить OpenAPI (Swagger UI) — теги: `auth`, `user`, `admin`, `asr`, `payments`, `api-keys`. +- Добавить примеры запросов/ответов для всех новых endpoints. + +--- + +## Этап 9. Деплой и окружение + +**Задача:** Подготовить к production. + +### Конфигурация +- Переменные окружения в `.env.example`: + - `DATABASE_URL` (PostgreSQL, async) + - `REDIS_URL` (rate limiting, refresh token blacklist, WS state кластеризации) + - `JWT_SECRET`, `JWT_ACCESS_EXPIRE_MINUTES`, `JWT_REFRESH_EXPIRE_DAYS` + - `ADMIN_SECRET_KEY` (для первичного создания superadmin через CLI) + - `YOOKASSA_SHOP_ID`, `YOOKASSA_SECRET_KEY`, `YOOKASSA_RETURN_URL` + - `RATE_LIMIT_STORAGE` (`redis` или `memory`) + - `MAX_FILE_SIZE_MB`, `ALLOWED_AUDIO_TYPES` + - `PROMETHEUS_ENABLED` (`true`/`false`) + - `TELEGRAM_BOT_TOKEN` (токен бота от @BotFather) + - `TELEGRAM_WEBAPP_URL` (публичный URL для Web App, например `https://asr.example.com/tg`) + - `TELEGRAM_PAYMENT_PROVIDER_TOKEN` (опционально, для `openInvoice` внутри TG) + +### Docker +- `Dockerfile` — многоэтапная сборка, установка Python-зависимостей, копирование Jinja-шаблонов и статики, `HEALTHCHECK`. +- `docker-compose.yml` — сервисы: `app`, `postgres`, `redis`, `nginx`. +- `docker-compose.override.yml` — для локальной разработки (volume mounts, hot reload). + +### Nginx +- Раздача статики (`static/`). +- Проксирование `/api/v1/*` и WebSocket `Upgrade`. +- `proxy_read_timeout` / `proxy_send_timeout` для WS (больше 60 сек). +- SSL/TLS (Let's Encrypt). + +### Безопасность +- **ORM:** защита от SQL-инъекций (SQLAlchemy). +- **XSS:** CSP заголовок, экранирование в Jinja2 (autoescape). +- **CSRF:** для форм использовать CSRF-токен в шаблоне (для JWT-форм). +- **Файлы:** валидация MIME-type и размера до чтения. +- **API Keys:** хранение только хешей, никогда не возвращать plain key после создания. +- **WS Auth:** обязательный auth фрейм, таймаут 5 сек на подключение без auth. + +--- + +## Этап 10. Документация и примеры + +**Задача:** Подготовить документацию для разработчиков и пользователей API. + +### Документация +- Обновить `README.md`: инструкции по запуску с Docker, переменные окружения, структура проекта. +- Создать `API_CHANGELOG.md` — изменения API версий. +- Добавить `docs/auth.md` — описание схемы JWT + API Key. +- Добавить `docs/payments.md` — интеграция с ЮKassa, обработка webhooks. + +### Примеры кода +- `examples/api_key_usage.py` — пример распознавания файла через API Key. +- `examples/websocket_client_auth.py` — пример WS-подключения с auth фреймом. +- Обновить `examples/streaming_client.py` для поддержки авторизации. + +### Telegram Web App документация +- `docs/telegram.md` — инструкция по настройке бота (@BotFather, Web App URL, HMAC-проверка). +- `examples/tg_webapp_auth.html` — минимальный пример HTML-страницы с авторизацией через `initData`. +- `examples/tg_bot_webhook.py` — пример FastAPI-обработчика webhook для Telegram-бота. + +--- + +## Чек-лист для aider + +Перед началом работы уточнить у пользователя: + +- [x] Какой ORM используется (SQLAlchemy, Tortoise, Prisma)? + *(Решение: SQLAlchemy 2.0 + Alembic. БД создаётся с нуля.)* +- [x] Есть ли Alembic/миграции? + *(Решение: нет, создаём с нуля.)* +- [x] Какой фронтенд-стек предпочтителен (чистый HTML+JS как в примерах, или Vue/React/HTMX)? + *(Решение: Jinja2 + чистый HTML+JS. Серверный рендеринг шаблонов.)* +- [x] Где хранить JWT — `localStorage` или `httpOnly` cookie? + *(Решение: Access токен — в памяти (memory), Refresh токен — `httpOnly` `Secure` `SameSite=Strict` cookie. WebSocket: первый фрейм после connect с токеном.)* +- [x] Нужна ли интеграция с конкретным платёжным провайдером (ЮKassa, Stripe) или пока заглушки? + *(Решение: интеграция с ЮKassa.)* +- [x] Rate limiting — по количеству запросов или по минутам аудио? + *(Решение: по количеству запросов.)* +- [x] Нужен ли API-доступ (API keys) или только веб-интерфейс? + *(Решение: нужен API-доступ через API keys.)* +- [x] Будет ли приложение доступно через Telegram Web App? + *(Решение: да, требуется интеграция с Telegram Bot API, авторизация через `initData`, отдельный мобильный UI.)* + +Общие для всех задач правила: +- Всегда комментируй код +- /no-commits +- Выполняя задачу, ставь отметку о её выполнении в этом файле. +- Выполняя задачу анализируй как изменение плана или результата выполнения влияет на другие задачи и корректируй их и останавливай выполнение задачи, если нужно решение по корректировке от пользователя. +- Выполняя задачу ты можешь запрашивать у пользователя необходимые файлы. +- Тебе нужно следить за соответствием моделей в alembic. +- **Гибридная модель:** HTMX для CRUD-страниц (история, профиль, ключи), чистый JS для WebSocket (ASR, мониторинг). +- **Единая дизайн-система:** CSS-переменные в `static/css/design-system.css`, общие для user и admin. +- **Состояния UI:** каждый интерактивный элемент имеет 4 состояния — idle, loading, success, error. Нет `alert()`. +- **CSS-изоляция:** пользовательские стили в `static/css/design-system.css`, админские доработки — там же через переменные. Не создавай конфликтов с существующими admin-стилями. +- **JS-изоляция:** глобальные функции только в `ui.js` (toast, confirm). Всё остальное — модули или IIFE. +- **WebSocket формат:** для всех новых WS-клиентов (admin dashboard, user ASR) использовать строго: + ```json + {"type":"auth","access_token":"..."} + ``` + сразу после `onopen`, в течение 5 секунд. +- **Не трогай существующие API endpoints** в `main.py` / роутерах. Только добавляй HTML-роуты. +- **Не меняй модели SQLAlchemy** (они уже готовы). Используй как есть. +- **Сохрани `templates/index.html`** (legacy demo) — не удаляй, он используется для `/demo`. +- **Сохрани `templates/admin/base_admin.html`** — дорабатывай, но не ломай существующую структуру sidebar. diff --git a/utils/chunk_doing.py b/utils/chunk_doing.py index 37652e0..0b41c90 100644 --- a/utils/chunk_doing.py +++ b/utils/chunk_doing.py @@ -176,104 +176,134 @@ async def find_last_speech_position_v2(session, is_last_chunk): Версия find_last_speech_position без глобальных dict. Работает с полями session: audio_buffer, audio_overlap, audio_to_asr. """ + import time + vad_start = time.perf_counter() + if is_last_chunk: last_audio = session.audio_overlap + session.audio_buffer for i in range(0, len(last_audio), settings.MAX_OVERLAP_DURATION * 1000): session.audio_to_asr.append( last_audio[i:min(i + settings.MAX_OVERLAP_DURATION * 1000, len(last_audio))] ) + logger.debug( + "VAD last_chunk for %s: split into %d segments, total=%.3f sec, elapsed=%.3f sec", + session.client_id, + len(session.audio_to_asr), + last_audio.duration_seconds, + time.perf_counter() - vad_start, + ) + return + + # --- Не последний чанк: ищем паузу --- + session.audio_buffer = session.audio_overlap + session.audio_buffer + frame_rate = session.audio_buffer.frame_rate + silero_bitrate = 16000 + + if not session.audio_buffer: + logger.error("Ошибка: audio_buffer пустой") + raise ValueError("audio_buffer не может быть пустым") + + if session.audio_buffer.frame_rate != silero_bitrate: + audio_for_vad = await async_resample_audiosegment(session.audio_buffer, silero_bitrate) else: - session.audio_buffer = session.audio_overlap + session.audio_buffer - frame_rate = session.audio_buffer.frame_rate - silero_bitrate = 16000 - - if not session.audio_buffer: - logger.error("Ошибка: audio_buffer пустой") - raise ValueError("audio_buffer не может быть пустым") - - if session.audio_buffer.frame_rate != silero_bitrate: - audio_for_vad = await async_resample_audiosegment(session.audio_buffer, silero_bitrate) - else: - audio_for_vad = session.audio_buffer - - logger.debug(f"Получено из буфера на обработку аудио продолжительностью {session.audio_buffer.duration_seconds}") - - audio_for_vad = session.audio_overlap + audio_for_vad - + audio_for_vad = session.audio_buffer + + logger.debug( + "VAD start for %s: buffer=%.3f sec, overlap=%.3f sec, frame_rate=%d", + session.client_id, + session.audio_buffer.duration_seconds, + session.audio_overlap.duration_seconds, + frame_rate, + ) + + # Добавляем overlap к VAD-аудио для поиска границы (как в оригинале) + audio_for_vad = session.audio_overlap + audio_for_vad + + try: + audio = get_np_array_samples_float32(audio_for_vad.raw_data, audio_for_vad.sample_width) + logger.debug(f"Аудио для VAD: длина={len(audio)}, min={np.min(audio)}, max={np.max(audio)}") + except Exception as e: + logger.error(f"Ошибка в get_np_array_samples_float32: {e}") + raise + + if np.any(np.isnan(audio)) or np.any(np.isinf(audio)): + logger.error("Обнаружены NaN или бесконечные значения в audio") + raise ValueError("Некорректные значения в audio") + + duration_seconds = 0.5 + frame_length = 512 if audio_for_vad.frame_rate == 16000 else 256 + + if frame_length is None: + raise ValueError("для VAD Поддерживаются только фреймрейты 8000 или 16000 Гц") + + frame_duration = frame_length / frame_rate + min_silence_frames = int(duration_seconds / frame_duration) + max_audio_length = len(audio) if len(audio) < settings.MAX_OVERLAP_DURATION * silero_bitrate else settings.MAX_OVERLAP_DURATION * silero_bitrate + partial_frame_length = 0 + + frames = [audio[i:i + frame_length] for i in range(int(len(audio) // 3), max_audio_length, frame_length)] + logger.debug(f"Создано фреймов: {len(frames)}, frame_length={frame_length}") + + silence_frames = 0 + await vad.reset_state() + vad_state = vad.state + + no_silent = False + for i, frame in enumerate(reversed(frames)): + vad.state = vad_state try: - audio = get_np_array_samples_float32(audio_for_vad.raw_data, audio_for_vad.sample_width) - logger.debug(f"Аудио для VAD: длина={len(audio)}, min={np.min(audio)}, max={np.max(audio)}") + if len(frame) < frame_length: + partial_frame_length = len(frame) + logger.debug(f"Пропущен неполный фрейм: длина={partial_frame_length}") + continue + else: + speech_prob, vad_state = await vad.is_speech(frame, audio_for_vad.frame_rate) + if speech_prob < vad.prob_level: + silence_frames += 1 + if silence_frames >= min_silence_frames: + break + else: + silence_frames = 0 except Exception as e: - logger.error(f"Ошибка в get_np_array_samples_float32: {e}") + logger.error(f"Ошибка VAD - {e}" + f"\nframe_rate = {frame_rate}" + f"\nframe_length = {frame_length}" + f"\nframe_index = {i}" + f"\nframe_length_actual = {len(frame)}") raise + else: + no_silent = True - if np.any(np.isnan(audio)) or np.any(np.isinf(audio)): - logger.error("Обнаружены NaN или бесконечные значения в audio") - raise ValueError("Некорректные значения в audio") - - duration_seconds = 0.5 - frame_length = 512 if audio_for_vad.frame_rate == 16000 else 256 - - if frame_length is None: - raise ValueError("для VAD Поддерживаются только фреймрейты 8000 или 16000 Гц") - - frame_duration = frame_length / frame_rate - min_silence_frames = int(duration_seconds / frame_duration) - max_audio_length = len(audio) if len(audio) < settings.MAX_OVERLAP_DURATION * silero_bitrate else settings.MAX_OVERLAP_DURATION * silero_bitrate - partial_frame_length = 0 - - frames = [audio[i:i + frame_length] for i in range(int(len(audio) // 3), max_audio_length, frame_length)] - logger.debug(f"Создано фреймов: {len(frames)}, frame_length={frame_length}") - - silence_frames = 0 - await vad.reset_state() - vad_state = vad.state - - no_silent = False - for i, frame in enumerate(reversed(frames)): - vad.state = vad_state - try: - if len(frame) < frame_length: - partial_frame_length = len(frame) - logger.debug(f"Пропущен неполный фрейм: длина={partial_frame_length}") - continue - else: - logger.debug(f"Обработка фрейма {i}: длина={len(frame)}, min={np.min(frame)}, max={np.max(frame)}") - speech_prob, vad_state = await vad.is_speech(frame, audio_for_vad.frame_rate) - if speech_prob < vad.prob_level: - logger.debug(f"Найден не голос на speech_end = {max_audio_length-(i+1)*frame_length-partial_frame_length}") - silence_frames += 1 - if silence_frames >= min_silence_frames: - break - else: - silence_frames = 0 - logger.debug(f"Найден ГОЛОС на speech_end = {max_audio_length-i*frame_length-partial_frame_length}") - except Exception as e: - logger.error(f"Ошибка VAD - {e}" - f"\nframe_rate = {frame_rate}" - f"\nframe_length = {frame_length}" - f"\nframe_index = {i}" - f"\nframe_length_actual = {len(frame)}") - raise - else: - no_silent = True - - try: - if no_silent: - speech_end = max_audio_length - elif not partial_frame_length: - speech_end = max_audio_length - (i + 1) * frame_length - else: - speech_end = max_audio_length - i * frame_length - except Exception as e: - print(e) + try: + if no_silent: + speech_end = max_audio_length + elif not partial_frame_length: + speech_end = max_audio_length - (i + 1) * frame_length else: - separation_time = int(speech_end * 1000 / silero_bitrate) - session.audio_to_asr.append(session.audio_buffer[:separation_time]) - session.audio_overlap = session.audio_buffer[separation_time:] - - logger.debug(f"Передано на ASR аудио продолжительностью {session.audio_to_asr[-1].duration_seconds}") - logger.debug(f"Передано в перекрытие аудио продолжительностью {session.audio_overlap.duration_seconds}") - session.audio_buffer = AudioSegment.silent(1, frame_rate) + speech_end = max_audio_length - i * frame_length + except Exception as e: + logger.error(f"Ошибка вычисления speech_end: {e}") + speech_end = max_audio_length + no_silent = True + + separation_time = int(speech_end * 1000 / silero_bitrate) + asr_segment = session.audio_buffer[:separation_time] + overlap_segment = session.audio_buffer[separation_time:] + + session.audio_to_asr.append(asr_segment) + session.audio_overlap = overlap_segment if overlap_segment.duration_seconds > 0 else AudioSegment.silent(1, frame_rate) + session.audio_buffer = AudioSegment.silent(1, frame_rate) + + logger.debug( + "VAD done for %s: no_silent=%s, speech_end=%d, separation_time=%d ms, " + "asr_segment=%.3f sec, overlap=%.3f sec, elapsed=%.3f sec", + session.client_id, + no_silent, + speech_end, + separation_time, + asr_segment.duration_seconds, + session.audio_overlap.duration_seconds, + time.perf_counter() - vad_start, + ) return From 4b76e5bcd9dda7b261bd20407019a4c8dc31d216 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Sun, 10 May 2026 12:16:46 +0300 Subject: [PATCH 33/46] =?UTF-8?q?=D0=94=D0=BE=D0=B1=D0=B0=D0=B2=D0=BB?= =?UTF-8?q?=D0=B5=D0=BD=D0=B0=20=D0=B0=D0=B4=D0=BC=D0=B8=D0=BD-=D0=BF?= =?UTF-8?q?=D0=B0=D0=BD=D0=B5=D0=BB=D1=8C.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/v1/endpoints/admin.py | 88 +++++++++++++++++--- api/v1/endpoints/admin_ws.py | 10 ++- core/deps.py | 4 +- core/exception_handlers.py | 11 ++- static/js/admin_dashboard.js | 139 ++++++++++++++++++++++++++++++++ static/js/admin_ui.js | 48 +++++++++++ static/js/auth.js | 21 ++++- templates/admin/base_admin.html | 97 ++++++++++++++++------ templates/admin/dashboard.html | 109 ++++++++++++++++++++++--- templates/admin/login.html | 3 +- 10 files changed, 474 insertions(+), 56 deletions(-) create mode 100644 static/js/admin_dashboard.js create mode 100644 static/js/admin_ui.js diff --git a/api/v1/endpoints/admin.py b/api/v1/endpoints/admin.py index 6ba2bf9..0dcf4bc 100644 --- a/api/v1/endpoints/admin.py +++ b/api/v1/endpoints/admin.py @@ -1,5 +1,7 @@ """FastAPI-роутер админ-панели.""" +from collections import defaultdict +from datetime import datetime, timedelta from typing import Optional from fastapi import APIRouter, Depends, HTTPException, status @@ -7,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from core.deps import get_current_user, require_admin, require_superadmin -from db.models import ApiKey, Plan, Subscription, Transaction, User +from db.models import ApiKey, Plan, Subscription, SystemLog, Transaction, User from db.session import get_db_session from models.admin import ( AdminApiKeyResponse, @@ -39,9 +41,75 @@ async def admin_metrics(current_user: User = Depends(require_admin)): @router.get("/metrics/history") -async def admin_metrics_history(current_user: User = Depends(require_admin)): - """История метрик за период (заглушка).""" - return {"detail": "История метрик — заглушка"} +async def admin_metrics_history( + range: str = "24h", + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """История метрик за период с группировкой по 5-минутным слотам.""" + # Парсим параметр range (поддерживаем 1h, 24h, 7d и т.д.) + try: + if range.endswith("h"): + hours = int(range[:-1]) + since = datetime.utcnow() - timedelta(hours=hours) + elif range.endswith("d"): + days = int(range[:-1]) + since = datetime.utcnow() - timedelta(days=days) + else: + since = datetime.utcnow() - timedelta(hours=24) + except ValueError: + since = datetime.utcnow() - timedelta(hours=24) + + result = await db.execute( + select(SystemLog) + .where( + SystemLog.component == "SystemMetricsCollector", + SystemLog.created_at >= since, + ) + .order_by(SystemLog.created_at) + ) + logs = result.scalars().all() + + # Группировка по 5-минутным слотам + slots = defaultdict( + lambda: { + "cpu_values": [], + "gpu_values": [], + "conn_values": [], + "queue_values": [], + } + ) + + for log in logs: + ts = log.created_at + slot_ts = ts.replace( + minute=(ts.minute // 5) * 5, second=0, microsecond=0 + ) + key = slot_ts.isoformat() + meta = log.meta or {} + slots[key]["cpu_values"].append(meta.get("cpu_percent")) + slots[key]["gpu_values"].append(meta.get("gpu_utilization_percent")) + slots[key]["conn_values"].append(meta.get("active_connections")) + slots[key]["queue_values"].append(meta.get("queue_depth")) + + def _avg(values): + clean = [v for v in values if v is not None] + return round(sum(clean) / len(clean), 2) if clean else None + + response = [] + for key in sorted(slots.keys()): + data = slots[key] + response.append( + { + "timestamp": key, + "cpu": _avg(data["cpu_values"]), + "gpu": _avg(data["gpu_values"]), + "active_connections": _avg(data["conn_values"]), + "queue_depth": _avg(data["queue_values"]), + } + ) + + return response @router.get("/users", response_model=list[AdminUserListItem]) @@ -67,7 +135,7 @@ async def admin_users_list( id=u.id, email=u.email, full_name=u.full_name, - role=u.role.value, + role=u.role, is_active=u.is_active, created_at=u.created_at, last_login_at=u.last_login_at, @@ -93,7 +161,7 @@ async def admin_user_detail( id=user.id, email=user.email, full_name=user.full_name, - role=user.role.value, + role=user.role, is_active=user.is_active, created_at=user.created_at, last_login_at=user.last_login_at, @@ -124,7 +192,7 @@ async def admin_user_update( id=user.id, email=user.email, full_name=user.full_name, - role=user.role.value, + role=user.role, is_active=user.is_active, created_at=user.created_at, last_login_at=user.last_login_at, @@ -306,7 +374,7 @@ async def admin_subscriptions_list( user_id=s.user_id, plan_id=s.plan_id, plan_name=None, - status=s.status.value, + status=s.status, started_at=s.started_at, expires_at=s.expires_at, auto_renew=s.auto_renew, @@ -367,7 +435,7 @@ async def admin_transactions_list( subscription_id=t.subscription_id, amount=float(t.amount) if t.amount is not None else None, currency=t.currency, - status=t.status.value, + status=t.status, payment_provider=t.payment_provider, external_payment_id=t.external_payment_id, created_at=t.created_at, @@ -395,7 +463,7 @@ async def admin_logs_list( return [ AdminSystemLogResponse( id=l.id, - level=l.level.value, + level=l.level, component=l.component, message=l.message, meta=l.meta, diff --git a/api/v1/endpoints/admin_ws.py b/api/v1/endpoints/admin_ws.py index 69e015f..592aa3b 100644 --- a/api/v1/endpoints/admin_ws.py +++ b/api/v1/endpoints/admin_ws.py @@ -6,6 +6,7 @@ from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status from sqlalchemy import select +from core.exceptions import InvalidTokenException, TokenExpiredException from core.security import decode_token # type: ignore[import-untyped] from db.enums import UserRole from db.models import User @@ -31,7 +32,11 @@ async def admin_websocket(websocket: WebSocket): return token = auth_msg.get("access_token", "") - payload = decode_token(token) + try: + payload = decode_token(token) + except (TokenExpiredException, InvalidTokenException): + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return if not payload or not getattr(payload, "sub", None): await websocket.close(code=status.WS_1008_POLICY_VIOLATION) return @@ -53,7 +58,8 @@ async def admin_websocket(websocket: WebSocket): active_connections=ws_manager.active_connections_count if ws_manager else 0, max_connections=getattr(ws_manager, "max_connections", 100), ) - await websocket.send_json({"type": "metrics", "data": metrics}) + metrics_data = metrics.model_dump() if hasattr(metrics, "model_dump") else metrics + await websocket.send_json({"type": "metrics", "data": metrics_data}) else: await websocket.send_json( {"type": "metrics", "data": {"detail": "Collector unavailable"}} diff --git a/core/deps.py b/core/deps.py index 497e57c..5fb8f07 100644 --- a/core/deps.py +++ b/core/deps.py @@ -60,7 +60,7 @@ async def get_current_user( def require_admin(user: User = Depends(get_current_user)) -> User: """Проверяет, что пользователь — админ или суперадмин.""" - if user.role not in (UserRole.admin, UserRole.superadmin): + if str(user.role) not in (UserRole.admin.value, UserRole.superadmin.value): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Требуются права администратора", @@ -70,7 +70,7 @@ def require_admin(user: User = Depends(get_current_user)) -> User: def require_superadmin(user: User = Depends(get_current_user)) -> User: """Проверяет, что пользователь — суперадмин.""" - if user.role != UserRole.superadmin: + if str(user.role) != UserRole.superadmin.value: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Требуются права суперадминистратора", diff --git a/core/exception_handlers.py b/core/exception_handlers.py index f1d1f41..bf7b818 100644 --- a/core/exception_handlers.py +++ b/core/exception_handlers.py @@ -3,7 +3,7 @@ from fastapi import Request from fastapi.exceptions import RequestValidationError -from fastapi.responses import JSONResponse +from fastapi.responses import JSONResponse, RedirectResponse from starlette.exceptions import HTTPException as StarletteHTTPException from models.fast_api_models import ErrorResponse @@ -23,8 +23,17 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE async def http_exception_handler(request: Request, exc: StarletteHTTPException): + accept = request.headers.get("accept", "") + is_html = "text/html" in accept + + if is_html and exc.status_code in (401, 403): + login_url = "/admin/login" if request.url.path.startswith("/admin") else "/login" + if request.url.path not in (login_url, "/login", "/admin/login"): + return RedirectResponse(url=login_url, status_code=303) + return JSONResponse( status_code=exc.status_code, + headers=exc.headers, content=ErrorResponse( success=False, error_description=exc.detail, diff --git a/static/js/admin_dashboard.js b/static/js/admin_dashboard.js new file mode 100644 index 0000000..4f6f9dc --- /dev/null +++ b/static/js/admin_dashboard.js @@ -0,0 +1,139 @@ +(function() { + 'use strict'; + + class AdminDashboardWS { + constructor() { + this.ws = null; + this.reconnectAttempts = 0; + this.maxReconnectAttempts = 10; + this.baseDelay = 1000; + this.maxDelay = 30000; + this.heartbeatInterval = null; + this._isConnected = false; + } + + connect(token) { + this._token = token; + this._connect(); + } + + _connect() { + if (this.ws) { + try { this.ws.close(); } catch(e) {} + } + + const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; + const url = `${protocol}//${window.location.host}/api/v1/admin/ws`; + + this.ws = new WebSocket(url); + + this.ws.onopen = () => { + this.reconnectAttempts = 0; + this._isConnected = true; + if (this._token) { + this.ws.send(JSON.stringify({ type: 'auth', access_token: this._token })); + } + this._startHeartbeat(); + this._setWidgetsStatus('connected'); + }; + + this.ws.onmessage = (event) => { + try { + const msg = JSON.parse(event.data); + if (msg.type === 'metrics') { + this._updateWidgets(msg.data); + } else if (msg.type === 'alert') { + this._showAlert(msg.data); + } + } catch (e) { + console.error('WS parse error', e); + } + }; + + this.ws.onerror = () => { + this._setWidgetsStatus('error'); + }; + + this.ws.onclose = (event) => { + this._isConnected = false; + this._stopHeartbeat(); + // Если закрытие из-за истёкшего токена — пробуем refresh и сразу переподключаемся + if (event.code === 1008 || event.code === 1011) { + Auth.refreshToken().then((ok) => { + if (ok) { + this.reconnectAttempts = 0; + this._connect(); + } else { + Auth.clearAuth(); + window.location.href = '/admin/login'; + } + }); + return; + } + if (this.reconnectAttempts < this.maxReconnectAttempts) { + const delay = Math.min(this.baseDelay * Math.pow(2, this.reconnectAttempts), this.maxDelay); + this.reconnectAttempts++; + setTimeout(() => this._connect(), delay); + this._setWidgetsStatus('reconnecting'); + } else { + this._setWidgetsStatus('disconnected'); + } + }; + } + + _startHeartbeat() { + this.heartbeatInterval = setInterval(() => { + if (this.ws && this.ws.readyState === WebSocket.OPEN) { + this.ws.send(JSON.stringify({ type: 'ping' })); + } + }, 30000); + } + + _stopHeartbeat() { + if (this.heartbeatInterval) { + clearInterval(this.heartbeatInterval); + this.heartbeatInterval = null; + } + } + + _updateWidgets(data) { + const setText = (id, text) => { + const el = document.getElementById(id); + if (el) el.innerHTML = `
${id.split('-')[1].toUpperCase()}
${text}
`; + }; + + setText('widget-cpu', data.cpu_utilization_percent != null ? data.cpu_utilization_percent.toFixed(1) + '%' : '—'); + setText('widget-gpu', data.gpu_utilization_percent != null ? data.gpu_utilization_percent.toFixed(1) + '%' : '—'); + setText('widget-tasks', data.active_tasks != null ? data.active_tasks : '—'); + setText('widget-queue', data.queue_depth != null ? data.queue_depth : '—'); + setText('widget-uptime', data.uptime_formatted || '—'); + setText('widget-adapter-status', data.adapter_status || '—'); + } + + _showAlert(data) { + const banner = document.getElementById('alertBanner'); + const text = document.getElementById('alertText'); + if (!banner || !text) return; + if (data.cpu_utilization_percent > 90 || data.queue_depth > 50) { + text.textContent = `CPU: ${data.cpu_utilization_percent?.toFixed(1)}%, Очередь: ${data.queue_depth}`; + banner.style.display = 'block'; + } + } + + _setWidgetsStatus(status) { + const color = status === 'connected' ? 'var(--color-success)' : status === 'reconnecting' ? 'var(--color-warning)' : 'var(--color-danger)'; + const label = status === 'connected' ? 'Live' : status === 'reconnecting' ? 'Reconnect...' : 'Offline'; + // Можно добавить индикатор статуса где-то на странице + } + + disconnect() { + this.reconnectAttempts = this.maxReconnectAttempts; + this._stopHeartbeat(); + if (this.ws) { + try { this.ws.close(1000, 'Client disconnect'); } catch(e) {} + } + } + } + + window.AdminDashboardWS = AdminDashboardWS; +})(); diff --git a/static/js/admin_ui.js b/static/js/admin_ui.js new file mode 100644 index 0000000..38511bd --- /dev/null +++ b/static/js/admin_ui.js @@ -0,0 +1,48 @@ +(function() { + 'use strict'; + + function renderBadge(status) { + const map = { + completed: 'badge-success', + active: 'badge-success', + failed: 'badge-danger', + cancelled: 'badge-danger', + processing: 'badge-warning', + pending: 'badge-warning', + expired: 'badge-warning', + }; + const cls = map[status] || 'badge-info'; + return `${status}`; + } + + function renderDate(iso) { + if (!iso) return '—'; + const d = new Date(iso); + const dd = String(d.getDate()).padStart(2, '0'); + const mm = String(d.getMonth() + 1).padStart(2, '0'); + const yyyy = d.getFullYear(); + const hh = String(d.getHours()).padStart(2, '0'); + const min = String(d.getMinutes()).padStart(2, '0'); + return `${dd}.${mm}.${yyyy} ${hh}:${min}`; + } + + function renderDuration(sec) { + if (!sec) return '0:00'; + const m = Math.floor(sec / 60); + const s = Math.floor(sec % 60); + return `${m}:${String(s).padStart(2, '0')}`; + } + + async function initAdmin() { + if (!Auth.getAccessToken()) { + const refreshed = await Auth.refreshToken(); + if (!refreshed) { + window.location.href = '/admin/login'; + return false; + } + } + return true; + } + + window.AdminUI = { renderBadge, renderDate, renderDuration, initAdmin }; +})(); diff --git a/static/js/auth.js b/static/js/auth.js index 3822878..ae0842f 100644 --- a/static/js/auth.js +++ b/static/js/auth.js @@ -1,9 +1,24 @@ (function() { + 'use strict'; + + // Восстанавливаем токен из cookie для SSR-переходов let _token = null; + const m = document.cookie.match(/(?:^|; )access_token=([^;]*)/); + if (m) { + try { _token = decodeURIComponent(m[1]); } catch(e) { _token = null; } + } + + function _setCookie(name, value, days) { + const expires = value + ? '; expires=' + new Date(Date.now() + days * 864e5).toUTCString() + : '; expires=Thu, 01 Jan 1970 00:00:00 GMT'; + const secure = window.location.protocol === 'https:' ? '; Secure' : ''; + document.cookie = name + '=' + encodeURIComponent(value || '') + expires + '; path=/; SameSite=Lax' + secure; + } - function setAccessToken(t) { _token = t; } + function setAccessToken(t) { _token = t; _setCookie('access_token', t, 1); } function getAccessToken() { return _token; } - function clearAuth() { _token = null; } + function clearAuth() { _token = null; _setCookie('access_token', '', -1); } let _isRefreshing = false; let _refreshPromise = null; @@ -85,5 +100,5 @@ } } - window.Auth = { setAccessToken, getAccessToken, clearAuth, apiFetch, initAuth }; + window.Auth = { setAccessToken, getAccessToken, clearAuth, apiFetch, initAuth, refreshToken }; })(); diff --git a/templates/admin/base_admin.html b/templates/admin/base_admin.html index ee7727a..7a25a14 100644 --- a/templates/admin/base_admin.html +++ b/templates/admin/base_admin.html @@ -1,35 +1,80 @@ - - - {% block title %}Админ-панель{% endblock %} - - + + + + {% block title %}Админ-панель{% endblock %} + + + + {% block head %}{% endblock %} -
+ +{% include "admin/components/tariff_modal.html" %} +{% endblock %} + +{% block scripts %} {% endblock %} From 8320b15f766d0f8d69eebb0ffc88da1a783c0c46 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Sun, 10 May 2026 15:04:30 +0300 Subject: [PATCH 35/46] =?UTF-8?q?=D0=94=D0=BE=D1=80=D0=B0=D0=B1=D0=BE?= =?UTF-8?q?=D1=82=D0=B0=D0=BD=D0=B0=20=D1=81=D1=82=D1=80=D0=B0=D0=BD=D0=B8?= =?UTF-8?q?=D1=86=D0=B0=20=D1=81=20=D1=81=D0=B5=D1=81=D1=81=D0=B8=D1=8F?= =?UTF-8?q?=D0=BC=D0=B8=20=D0=BF=D0=BE=D0=BB=D1=8C=D0=B7=D0=BE=D0=B2=D0=B0?= =?UTF-8?q?=D1=82=D0=B5=D0=BB=D0=B5=D0=B9=20=D0=B2=20=D0=B0=D0=B4=D0=BC?= =?UTF-8?q?=D0=B8=D0=BD=20=D0=BF=D0=B0=D0=BD=D0=B5=D0=BB=D0=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/v1/endpoints/admin.py | 57 ++++- api/v1/endpoints/auth.py | 2 +- routes/admin.py | 20 ++ templates/admin/session_detail.html | 130 +++++++++++ templates/admin/user_detail.html | 331 ++++++++++++++++++++++++++++ templates/admin/users.html | 148 ++++++++++--- 6 files changed, 655 insertions(+), 33 deletions(-) create mode 100644 templates/admin/session_detail.html create mode 100644 templates/admin/user_detail.html diff --git a/api/v1/endpoints/admin.py b/api/v1/endpoints/admin.py index 0dcf4bc..ed2042b 100644 --- a/api/v1/endpoints/admin.py +++ b/api/v1/endpoints/admin.py @@ -9,7 +9,8 @@ from sqlalchemy.ext.asyncio import AsyncSession from core.deps import get_current_user, require_admin, require_superadmin -from db.models import ApiKey, Plan, Subscription, SystemLog, Transaction, User +from core.security import create_access_token # type: ignore[import-untyped] +from db.models import ApiKey, ASRSession, Plan, Subscription, SystemLog, Transaction, User from db.session import get_db_session from models.admin import ( AdminApiKeyResponse, @@ -232,8 +233,8 @@ async def admin_user_sessions( return [ { "id": s.id, - "session_type": s.session_type.value, - "status": s.status.value, + "session_type": s.session_type, + "status": s.status, "audio_duration_sec": s.audio_duration_sec, "created_at": s.created_at, "completed_at": s.completed_at, @@ -254,10 +255,40 @@ async def admin_user_impersonate( raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail="Пользователь не найден" ) - token = await admin_service.impersonate_user(user) + token = create_access_token({"sub": user.id, "role": user.role}) return {"access_token": token, "token_type": "bearer"} +@router.get("/sessions/{session_id}") +async def admin_session_detail( + session_id: str, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Детали ASR-сессии.""" + result = await db.execute(select(ASRSession).where(ASRSession.id == session_id)) + s = result.scalar_one_or_none() + if not s: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="Сессия не найдена" + ) + return { + "id": s.id, + "user_id": s.user_id, + "session_type": s.session_type, + "status": s.status, + "audio_duration_sec": s.audio_duration_sec, + "processing_duration_sec": s.processing_duration_sec, + "cost": float(s.cost) if s.cost is not None else None, + "result_json": s.result_json, + "created_at": s.created_at, + "completed_at": s.completed_at, + "error_message": s.error_message, + "request_ip": s.request_ip, + "user_agent": s.user_agent, + } + + @router.get("/tariffs", response_model=list[AdminPlanResponse]) async def admin_tariffs_list( current_user: User = Depends(require_admin), @@ -361,13 +392,25 @@ async def admin_tariff_delete( @router.get("/subscriptions", response_model=list[AdminSubscriptionResponse]) async def admin_subscriptions_list( pagination: PaginationParams = Depends(), + user_id: Optional[str] = None, current_user: User = Depends(require_admin), db: AsyncSession = Depends(get_db_session), ): """Список подписок.""" - subs, total = await admin_service.get_subscriptions( - db, page=pagination.page, per_page=pagination.per_page - ) + if user_id: + from sqlalchemy import select + result = await db.execute( + select(Subscription) + .where(Subscription.user_id == user_id) + .order_by(Subscription.created_at.desc()) + .offset((pagination.page - 1) * pagination.per_page) + .limit(pagination.per_page) + ) + subs = result.scalars().all() + else: + subs, total = await admin_service.get_subscriptions( + db, page=pagination.page, per_page=pagination.per_page + ) return [ AdminSubscriptionResponse( id=s.id, diff --git a/api/v1/endpoints/auth.py b/api/v1/endpoints/auth.py index 9907942..f1f1da1 100644 --- a/api/v1/endpoints/auth.py +++ b/api/v1/endpoints/auth.py @@ -6,7 +6,7 @@ from config import settings from core.deps import get_current_user -from core.security import decode_token # type: ignore[import-untyped] +from core.security import create_access_token, decode_token # type: ignore[import-untyped] from db.models import User from db.session import get_db_session from models.auth import ( diff --git a/routes/admin.py b/routes/admin.py index d771a8e..4bf01da 100644 --- a/routes/admin.py +++ b/routes/admin.py @@ -36,6 +36,16 @@ async def admin_users_page( return templates.TemplateResponse("admin/users.html", {"request": request}) +@router.get("/users/{user_id}") +async def admin_user_detail_page( + request: Request, + user_id: str, + _admin=Depends(require_admin), +): + """Страница деталей пользователя.""" + return templates.TemplateResponse("admin/user_detail.html", {"request": request, "user_id": user_id}) + + @router.get("/sessions") async def admin_sessions_page( request: Request, @@ -45,6 +55,16 @@ async def admin_sessions_page( return templates.TemplateResponse("admin/sessions.html", {"request": request}) +@router.get("/sessions/{session_id}") +async def admin_session_detail_page( + request: Request, + session_id: str, + _admin=Depends(require_admin), +): + """Страница деталей сессии.""" + return templates.TemplateResponse("admin/session_detail.html", {"request": request, "session_id": session_id}) + + @router.get("/subscriptions") async def admin_subscriptions_page( request: Request, diff --git a/templates/admin/session_detail.html b/templates/admin/session_detail.html new file mode 100644 index 0000000..0f34ee2 --- /dev/null +++ b/templates/admin/session_detail.html @@ -0,0 +1,130 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Сессия — Админ-панель{% endblock %} +{% block content %} +
+
+ ← Назад +

Детали сессии

+ +
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+
+
+
+ +
+
+

Результат

+ +
+
+
+
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/user_detail.html b/templates/admin/user_detail.html new file mode 100644 index 0000000..ead687e --- /dev/null +++ b/templates/admin/user_detail.html @@ -0,0 +1,331 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Пользователь — Админ-панель{% endblock %} +{% block content %} +
+
+ ← Назад +

Пользователь

+ +
+ +
+
+
+ +
+
+
+ + +
+
+ + +
+
+ + +
+
+ +
+
+
+ +
+
+
+ +
+
+
+ +
+ + + +
+
+ +
+
+ + +
+ +
+ + + + + + + + + + + + + + +
IDТипСтатусДлительностьСозданаЗавершена
+
+
+
+ +
+ + +
+
+ + +{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/users.html b/templates/admin/users.html index d083f4c..22fbfb5 100644 --- a/templates/admin/users.html +++ b/templates/admin/users.html @@ -1,34 +1,132 @@ {% extends "admin/base_admin.html" %} {% block title %}Пользователи{% endblock %} {% block content %} -

Пользователи

-
-

Загрузка списка пользователей...

- - - - - +
+
+

Пользователи

+
+ +
+
+ +
+ + + + +
+ + + +
+
+ +
+
IDEmailИмяРольАктивенTelegramДействия
+ + + + + + + + + + + + + +
EmailИмяРольСтатусTelegramРегистрацияПоследний вход
+
+
+
+
+ +
+ + +{% endblock %} + +{% block scripts %} {% endblock %} From 6814751ec3d172a9070b3bded331147639574a59 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Sun, 10 May 2026 20:33:21 +0300 Subject: [PATCH 36/46] =?UTF-8?q?=D1=80=D0=B5=D0=B6=D0=B8=D0=BC=20=D1=8D?= =?UTF-8?q?=D0=BA=D1=81=D0=BF=D0=B5=D1=80=D1=82=D0=B0=20=D0=BF=D1=80=D0=B8?= =?UTF-8?q?=20=D0=B2=D1=8B=D0=B1=D0=BE=D1=80=D0=B5=20=D0=BF=D0=B0=D1=80?= =?UTF-8?q?=D0=B0=D0=BC=D0=B5=D1=82=D1=80=D0=BE=D0=B2=20=D1=80=D0=B0=D1=81?= =?UTF-8?q?=D0=BF=D0=BE=D0=B7=D0=BD=D0=B0=D0=B2=D0=B0=D0=BD=D0=B8=D1=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- static/js/asr_page.js | 202 +++++++++++++++++--- static/js/asr_settings.js | 170 +++++++++++++++++ templates/user/asr.html | 387 ++++++++++++++++++++++++++++++++++---- to_do_front.txt | 213 ++++++++++++++++----- 4 files changed, 865 insertions(+), 107 deletions(-) create mode 100644 static/js/asr_settings.js diff --git a/static/js/asr_page.js b/static/js/asr_page.js index 23a9a76..b18fa8d 100644 --- a/static/js/asr_page.js +++ b/static/js/asr_page.js @@ -1,6 +1,22 @@ (function() { 'use strict'; + function renderAsrResult(payload, mode) { + const settings = ASRSettings.getFor(mode); + const sentenced = payload.sentenced_data || payload.data?.sentenced_data; + if (settings.split_phrases && sentenced?.list_of_sentenced_recognitions) { + return sentenced.list_of_sentenced_recognitions.map(item => { + const start = item.start != null ? item.start : (item.start_time || ''); + return `${start} - ${item.text || ''}`; + }).join('\n'); + } + if (!settings.split_phrases && sentenced?.full_text_only) { + const ft = sentenced.full_text_only; + return Array.isArray(ft) ? ft.join('\n') : String(ft); + } + return sentenced?.raw_text_sentenced_recognition || payload.data?.text || payload.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || JSON.stringify(payload, null, 2); + } + // --- Табы --- function switchTab(name) { document.querySelectorAll('.tab-panel').forEach(p => p.style.display = 'none'); @@ -28,28 +44,86 @@ window.downloadText = downloadText; // --- URL --- + function validateExpert(mode) { + let ok = true; + const showErr = (id, msg) => { + const el = document.getElementById(mode + '_err_' + id); + if (el) { el.textContent = msg; el.style.display = msg ? 'block' : 'none'; } + }; + if (mode !== 'ws') { + const sensity = parseInt(document.getElementById(mode + '_diar_vad_sensity').value); + if (isNaN(sensity) || sensity < 1 || sensity > 5) { + showErr('diar_vad_sensity', 'Допустимые значения: 1–5'); ok = false; + } else { showErr('diar_vad_sensity', ''); } + + const speed = parseFloat(document.getElementById(mode + '_speech_speed_correction_multiplier').value); + if (isNaN(speed) || speed <= 0) { + showErr('speech_speed_correction_multiplier', 'Должно быть > 0'); ok = false; + } else { showErr('speech_speed_correction_multiplier', ''); } + + const batch = parseInt(document.getElementById(mode + '_batch_size').value); + if (isNaN(batch) || batch < 1) { + showErr('batch_size', 'Минимум 1'); ok = false; + } else { showErr('batch_size', ''); } + } else { + const sampleRate = parseInt(document.getElementById('ws_sample_rate').value); + if (isNaN(sampleRate) || sampleRate < 8000 || sampleRate > 48000) { + showErr('sample_rate', 'Допустимый диапазон: 8000–48000'); ok = false; + } else { showErr('sample_rate', ''); } + } + return ok; + } + window.validateExpert = validateExpert; + + function getUrlParams() { + const expert = document.getElementById('url_expert').checked; + const defaults = ASRSettings.getDefaults('url'); + const domBool = (id) => document.getElementById(id).checked; + const domInt = (id, def) => { + const v = parseInt(document.getElementById(id).value); + return isNaN(v) ? def : v; + }; + const domFloat = (id, def) => { + const v = parseFloat(document.getElementById(id).value); + return isNaN(v) ? def : v; + }; + const fastSpeech = domBool('url_fast_speech'); + const splitPhrases = domBool('url_split_phrases'); + return { + AudioFileUrl: document.getElementById('urlInput').value.trim(), + keep_raw: expert ? domBool('url_keep_raw') : defaults.keep_raw, + do_echo_clearing: expert ? domBool('url_do_echo_clearing') : defaults.do_echo_clearing, + do_dialogue: splitPhrases ? true : (expert ? domBool('url_do_dialogue') : defaults.do_dialogue), + do_punctuation: splitPhrases ? true : (expert ? domBool('url_do_punctuation') : defaults.do_punctuation), + do_diarization: expert ? domBool('url_do_diarization') : defaults.do_diarization, + make_mono: expert ? domBool('url_make_mono') : defaults.make_mono, + diar_vad_sensity: expert ? domInt('url_diar_vad_sensity', defaults.diar_vad_sensity) : defaults.diar_vad_sensity, + do_auto_speech_speed_correction: fastSpeech ? true : (expert ? domBool('url_do_auto_speech_speed_correction') : defaults.do_auto_speech_speed_correction), + speech_speed_correction_multiplier: expert ? domFloat('url_speech_speed_correction_multiplier', defaults.speech_speed_correction_multiplier) : defaults.speech_speed_correction_multiplier, + use_batch: fastSpeech ? false : (expert ? domBool('url_use_batch') : defaults.use_batch), + batch_size: expert ? domInt('url_batch_size', defaults.batch_size) : defaults.batch_size, + }; + } + async function sendUrl() { const btn = document.getElementById('btnUrl'); UI.setLoading(btn, true); - const url = document.getElementById('urlInput').value.trim(); - if (!url) { UI.toast('Введите ссылку', 'warning'); UI.setLoading(btn, false); return; } + const payload = getUrlParams(); + if (!payload.AudioFileUrl) { UI.toast('Введите ссылку', 'warning'); UI.setLoading(btn, false); return; } + if (document.getElementById('url_expert').checked && !validateExpert('url')) { + UI.setLoading(btn, false); return; + } try { const resp = await Auth.apiFetch('/api/v1/asr/url', { method: 'POST', headers: {'Content-Type': 'application/json'}, - body: JSON.stringify({ - AudioFileUrl: url, - keep_raw: document.getElementById('url_keep_raw').checked, - do_echo_clearing: document.getElementById('url_do_echo').checked, - do_dialogue: document.getElementById('url_do_dialogue').checked, - do_punctuation: document.getElementById('url_do_punct').checked - }) + body: JSON.stringify(payload) }); const data = await resp.json(); const el = document.getElementById('urlResult'); const pre = document.getElementById('urlResultText'); el.style.display = 'block'; - pre.textContent = data.data?.sentenced_data?.raw_text_sentenced_recognition || data.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || JSON.stringify(data, null, 2); + pre.textContent = renderAsrResult(data, 'url'); } catch (e) { UI.toast('Ошибка: ' + e.message, 'error'); } finally { @@ -75,22 +149,50 @@ } window.handleFileSelect = handleFileSelect; + function getFileParams() { + const expert = document.getElementById('file_expert').checked; + const defaults = ASRSettings.getDefaults('file'); + const domBool = (id) => document.getElementById(id).checked; + const domInt = (id, def) => { + const v = parseInt(document.getElementById(id).value); + return isNaN(v) ? def : v; + }; + const domFloat = (id, def) => { + const v = parseFloat(document.getElementById(id).value); + return isNaN(v) ? def : v; + }; + const fastSpeech = domBool('file_fast_speech'); + const splitPhrases = domBool('file_split_phrases'); + const form = new FormData(); + form.append('keep_raw', expert ? domBool('file_keep_raw') : defaults.keep_raw); + form.append('do_echo_clearing', expert ? domBool('file_do_echo_clearing') : defaults.do_echo_clearing); + form.append('do_dialogue', splitPhrases ? true : (expert ? domBool('file_do_dialogue') : defaults.do_dialogue)); + form.append('do_diarization', expert ? domBool('file_do_diarization') : defaults.do_diarization); + form.append('do_punctuation', splitPhrases ? true : (expert ? domBool('file_do_punctuation') : defaults.do_punctuation)); + form.append('make_mono', expert ? domBool('file_make_mono') : defaults.make_mono); + form.append('diar_vad_sensity', expert ? domInt('file_diar_vad_sensity', defaults.diar_vad_sensity) : defaults.diar_vad_sensity); + form.append('do_auto_speech_speed_correction', fastSpeech ? true : (expert ? domBool('file_do_auto_speech_speed_correction') : defaults.do_auto_speech_speed_correction)); + form.append('speech_speed_correction_multiplier', expert ? domFloat('file_speech_speed_correction_multiplier', defaults.speech_speed_correction_multiplier) : defaults.speech_speed_correction_multiplier); + form.append('use_batch', fastSpeech ? false : (expert ? domBool('file_use_batch') : defaults.use_batch)); + form.append('batch_size', expert ? domInt('file_batch_size', defaults.batch_size) : defaults.batch_size); + return form; + } + async function sendFile() { if (!selectedFile) return; const btn = document.getElementById('btnFile'); UI.setLoading(btn, true); document.getElementById('fileProgress').style.display = 'block'; - const form = new FormData(); + if (document.getElementById('file_expert').checked && !validateExpert('file')) { + UI.setLoading(btn, false); document.getElementById('fileProgress').style.display = 'none'; return; + } + const form = getFileParams(); form.append('file', selectedFile); - form.append('keep_raw', document.getElementById('file_keep_raw').checked); - form.append('do_dialogue', document.getElementById('file_do_dialogue').checked); - form.append('do_diarization', document.getElementById('file_do_diar').checked); - form.append('do_punctuation', document.getElementById('file_do_punct').checked); try { const resp = await Auth.apiFetch('/api/v1/asr/file', {method:'POST', body:form}); const data = await resp.json(); document.getElementById('fileResult').style.display = 'block'; - document.getElementById('fileResultText').textContent = data.data?.sentenced_data?.raw_text_sentenced_recognition || data.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || JSON.stringify(data, null, 2); + document.getElementById('fileResultText').textContent = renderAsrResult(data, 'file'); } catch (e) { UI.toast('Ошибка: ' + e.message, 'error'); } finally { @@ -140,31 +242,54 @@ return window.btoa(binary); } - async function sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate) { + function getWsParams() { + const expert = document.getElementById('ws_expert').checked; + const defaults = ASRSettings.getDefaults('ws'); + const domBool = (id) => document.getElementById(id).checked; + const domInt = (id, def) => { + const v = parseInt(document.getElementById(id).value); + return isNaN(v) ? def : v; + }; + const domString = (id, def) => { + const el = document.getElementById(id); + return el ? (el.value || def) : def; + }; + const fastSpeech = domBool('ws_fast_speech'); + const splitPhrases = domBool('ws_split_phrases'); + return { + sample_rate: defaults.sample_rate, + audio_format: expert ? domString('ws_audio_format', defaults.audio_format) : defaults.audio_format, + audio_transport: expert ? domString('ws_audio_transport', defaults.audio_transport) : defaults.audio_transport, + wait_null_answers: expert ? domBool('ws_wait_null_answers') : defaults.wait_null_answers, + do_dialogue: splitPhrases ? true : (expert ? domBool('ws_do_dialogue') : defaults.do_dialogue), + do_punctuation: splitPhrases ? true : (expert ? domBool('ws_do_punctuation') : defaults.do_punctuation), + channel_name: expert ? (domString('ws_channel_name', defaults.channel_name) || null) : defaults.channel_name, + }; + } + + async function sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate, wsParams) { return new Promise((resolve, reject) => { const wsProtocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:'; const wsUrl = `${wsProtocol}//${window.location.host}/api/v1/asr/ws`; const socket = new WebSocket(wsUrl); wsSockets.push(socket); - const doDialogue = document.getElementById('ws_do_dialogue').checked; - const doPunctuation = document.getElementById('ws_do_punct').checked; const token = Auth.getAccessToken(); socket.onopen = async function() { - const useBase64 = document.getElementById('ws_use_base64').checked; + const useBase64 = wsParams.audio_transport === 'json_base64'; if (token) { socket.send(JSON.stringify({ type: 'auth', access_token: token })); } socket.send(JSON.stringify({ type: 'config', sample_rate: sampleRate, - audio_format: 'pcm16', - audio_transport: useBase64 ? 'json_base64' : 'binary', - wait_null_answers: false, - do_dialogue: doDialogue, - do_punctuation: doPunctuation, - channel_name: 'channel_' + (channel + 1) + audio_format: wsParams.audio_format, + audio_transport: wsParams.audio_transport, + wait_null_answers: wsParams.wait_null_answers, + do_dialogue: wsParams.do_dialogue, + do_punctuation: wsParams.do_punctuation, + channel_name: wsParams.channel_name || ('channel_' + (channel + 1)) })); const dataView = new Uint8Array(arrayBuffer); @@ -206,7 +331,7 @@ document.getElementById('wsPartial').textContent = `[Канал ${channel+1}]: ${text} ...`; } if (data.type === 'final_result' || data.last_message) { - const text = data.sentenced_data?.raw_text_sentenced_recognition || data.data?.text || JSON.stringify(data, null, 2); + const text = renderAsrResult(data, 'ws'); const current = document.getElementById('wsResultText').textContent; document.getElementById('wsResultText').textContent = current + (current ? '\n\n' : '') + `=== Канал ${channel+1} ===\n${text}`; document.getElementById('wsPartial').textContent = ''; @@ -228,9 +353,9 @@ }); } - async function sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate) { + async function sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate, wsParams) { for (let channel = 0; channel < numChannels; channel++) { - await sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate); + await sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate, wsParams); console.log("конец канала"); } } @@ -238,6 +363,9 @@ async function sendWs() { const file = document.getElementById('wsFileInput').files[0]; if (!file) return; + if (document.getElementById('ws_expert').checked && !validateExpert('ws')) { + return; + } UI.setLoading(document.getElementById('btnWs'), true); document.getElementById('wsResult').style.display = 'block'; document.getElementById('wsResultText').textContent = ''; @@ -245,12 +373,13 @@ document.getElementById('btnWsStop').style.display = 'inline-flex'; setWsStatus('connecting'); wsSockets = []; + const wsParams = getWsParams(); try { const arrayBuffer = await file.arrayBuffer(); const { sampleRate, numChannels } = readWavHeader(arrayBuffer); const chunkSize = 65536; - await sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate); + await sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate, wsParams); setWsStatus('disconnected'); UI.toast('Распознавание завершено', 'success'); } catch (e) { @@ -291,6 +420,19 @@ } } + // Live validation for expert number inputs + ['url','file','ws'].forEach(mode => { + const panel = document.getElementById(mode + '_expert_panel'); + if (!panel) return; + panel.querySelectorAll('input[type="number"]').forEach(input => { + input.addEventListener('input', () => { + if (document.getElementById(mode + '_expert').checked) { + validateExpert(mode); + } + }); + }); + }); + // Инициализация loadHistory(); })(); diff --git a/static/js/asr_settings.js b/static/js/asr_settings.js new file mode 100644 index 0000000..e0a0204 --- /dev/null +++ b/static/js/asr_settings.js @@ -0,0 +1,170 @@ +/** + * ASRSettings — модуль для хранения эксперт-настроек ASR в localStorage. + * Ключ: asr_expert_settings_v1 + * Поддерживает миграцию схемы и сброс к дефолтам. + */ +(function() { + 'use strict'; + + const STORAGE_KEY = 'asr_expert_settings_v1'; + const SCHEMA_VERSION = 1; + + // Дефолты по Pydantic-моделям (SyncASRRequest, PostFileRequest, WSConfigMessage) + const DEFAULTS = { + url: { + keep_raw: true, + do_echo_clearing: true, + do_dialogue: true, + do_punctuation: true, + do_diarization: false, + make_mono: true, + diar_vad_sensity: 3, + do_auto_speech_speed_correction: false, + speech_speed_correction_multiplier: 1.0, + use_batch: true, + batch_size: 8, + expert: false, + fast_speech: false, + split_phrases: true + }, + file: { + keep_raw: true, + do_echo_clearing: false, + do_dialogue: true, + do_punctuation: true, + do_diarization: false, + make_mono: true, + diar_vad_sensity: 3, + do_auto_speech_speed_correction: false, + speech_speed_correction_multiplier: 1.0, + use_batch: true, + batch_size: 8, + expert: false, + fast_speech: false, + split_phrases: true + }, + ws: { + sample_rate: 16000, + audio_format: 'pcm16', + audio_transport: 'json_base64', + wait_null_answers: true, + do_dialogue: true, + do_punctuation: true, + channel_name: null, + use_base64: true, + expert: false, + fast_speech: false, + split_phrases: true + } + }; + + function _loadRaw() { + try { + const raw = localStorage.getItem(STORAGE_KEY); + return raw ? JSON.parse(raw) : null; + } catch (e) { + console.warn('ASRSettings: failed to parse localStorage, resetting'); + return null; + } + } + + function _saveRaw(data) { + localStorage.setItem(STORAGE_KEY, JSON.stringify(data)); + } + + function _ensureStorage() { + let data = _loadRaw(); + if (!data || data._schema !== SCHEMA_VERSION) { + data = { _schema: SCHEMA_VERSION }; + Object.keys(DEFAULTS).forEach(mode => { + data[mode] = { ...DEFAULTS[mode] }; + }); + _saveRaw(data); + } + // Гарантируем наличие всех полей (если схема менялась частично) + Object.keys(DEFAULTS).forEach(mode => { + data[mode] = data[mode] || {}; + Object.keys(DEFAULTS[mode]).forEach(key => { + if (!(key in data[mode])) { + data[mode][key] = DEFAULTS[mode][key]; + } + }); + }); + return data; + } + + const ASRSettings = { + /** + * Вернуть настройки для режима (url | file | ws). + * Всегда возвращает полный объект с дефолтами для отсутствующих ключей. + */ + getFor(mode) { + const data = _ensureStorage(); + return { ...DEFAULTS[mode], ...(data[mode] || {}) }; + }, + + /** + * Сохранить настройки для режима. + * @param {string} mode — 'url', 'file', 'ws' + * @param {object} values — объект с обновлёнными значениями + */ + setFor(mode, values) { + const data = _ensureStorage(); + data[mode] = { ...(data[mode] || {}), ...values }; + _saveRaw(data); + }, + + /** + * Сбросить настройки режима к дефолтам. + */ + reset(mode) { + const data = _ensureStorage(); + data[mode] = { ...DEFAULTS[mode] }; + _saveRaw(data); + }, + + /** + * Полный сброс всех настроек. + */ + resetAll() { + const data = { _schema: SCHEMA_VERSION }; + Object.keys(DEFAULTS).forEach(mode => { + data[mode] = { ...DEFAULTS[mode] }; + }); + _saveRaw(data); + }, + + /** + * Проверить и при необходимости мигрировать схему. + * При несовпадении версии — сброс с уведомлением через callback. + */ + migrate(onReset) { + const data = _loadRaw(); + if (!data || data._schema !== SCHEMA_VERSION) { + this.resetAll(); + if (typeof onReset === 'function') onReset(); + return false; + } + return true; + }, + + /** + * Проверить, отличаются ли текущие настройки режима от дефолтов. + */ + isDirty(mode) { + const current = this.getFor(mode); + const defs = DEFAULTS[mode]; + return Object.keys(defs).some(key => current[key] !== defs[key]); + }, + + /** + * Вернуть дефолтные настройки для режима (копия). + */ + getDefaults(mode) { + return { ...DEFAULTS[mode] }; + } + }; + + // Экспорт в глобальную область + window.ASRSettings = ASRSettings; +})(); diff --git a/templates/user/asr.html b/templates/user/asr.html index fd8a50d..680db41 100644 --- a/templates/user/asr.html +++ b/templates/user/asr.html @@ -31,20 +31,71 @@

Распознавание речи

-
- - - -
+
+ + +
+
+ + +
+
+ + +
@@ -46,7 +58,7 @@

История сессий

{% endblock %} diff --git a/templates/user/profile.html b/templates/user/profile.html index f30af3f..15ab7b6 100644 --- a/templates/user/profile.html +++ b/templates/user/profile.html @@ -34,7 +34,7 @@

Смена пароля

- +
@@ -114,7 +114,28 @@

Опасная зона { + const dialog = document.createElement('dialog'); + dialog.innerHTML = ` +
+

Удалить аккаунт

+

Введите слово УДАЛИТЬ для подтверждения.

+ +
+ + +
+
+ `; + document.body.appendChild(dialog); + dialog.showModal(); + const input = dialog.querySelector('#confirmDeleteInput'); + const confirmBtn = dialog.querySelector('#dlg-confirm'); + input.addEventListener('input', () => { + confirmBtn.disabled = input.value.trim() !== 'УДАЛИТЬ'; + }); + dialog.querySelector('#dlg-cancel').onclick = () => { dialog.close(); dialog.remove(); }; + dialog.querySelector('#dlg-confirm').onclick = async () => { + dialog.close(); dialog.remove(); try { const resp = await Auth.apiFetch('/api/v1/user/profile', {method: 'DELETE'}); if (!resp.ok) throw new Error('Ошибка'); @@ -123,7 +144,8 @@

Опасная зона dialog.remove()); } loadProfile(); diff --git a/templates/user/subscription.html b/templates/user/subscription.html index cc3d1c1..446031f 100644 --- a/templates/user/subscription.html +++ b/templates/user/subscription.html @@ -13,17 +13,14 @@

Управление подпиской

Доступные тарифы

-
-
- Выбор и оплата тарифов — в разработке.
- Перейти к распознаванию -
+
+

История платежей

-
-
Нет данных
+
+
@@ -64,7 +61,7 @@

${s.plan_name || 'Тариф'}

- +
`; } catch (e) { @@ -73,8 +70,9 @@

${s.plan_name || 'Тариф'}

} } - async function cancelRenew() { + async function cancelRenew(btn) { UI.confirmDialog('Отключить авто-продление?', 'Вы уверены?', async () => { + if (btn) UI.setLoading(btn, true); try { const resp = await Auth.apiFetch('/api/v1/user/subscription/cancel', {method:'POST'}); if (!resp.ok) throw new Error('Ошибка'); @@ -82,10 +80,61 @@

${s.plan_name || 'Тариф'}

loadSubscription(); } catch (e) { UI.toast(e.message, 'error'); + } finally { + if (btn) UI.setLoading(btn, false); } }); } + async function loadTariffs() { + const el = document.getElementById('tariffsContainer'); + setTimeout(() => { + el.innerHTML = ` +
+ Выбор и оплата тарифов — в разработке.
+ Перейти к распознаванию +
+ `; + }, 500); + } + + async function loadPayments() { + const el = document.getElementById('paymentsContainer'); + try { + const resp = await Auth.apiFetch('/api/v1/user/payments'); + if (!resp.ok) throw new Error(); + const data = await resp.json(); + if (!data || !data.length) { + el.innerHTML = '
Нет платежей
'; + return; + } + el.innerHTML = ` + + + + + + + + + + ${data.map(p => ` + + + + + + `).join('')} + +
ДатаСуммаСтатус
${UI.formatDate(p.created_at)}${p.amount} ${p.currency}${p.status}
+ `; + } catch (e) { + el.innerHTML = '
Нет данных
'; + } + } + loadSubscription(); + loadTariffs(); + loadPayments(); {% endblock %} From 11ffede8a380f51128850e57ee7430547a53d5a3 Mon Sep 17 00:00:00 2001 From: Sanich137 Date: Mon, 11 May 2026 11:59:13 +0300 Subject: [PATCH 38/46] =?UTF-8?q?=D0=A3=D1=81=D1=82=D1=80=D0=B0=D0=BD?= =?UTF-8?q?=D0=B5=D0=BD=D1=8B=20=D0=B1=D0=B0=D0=B3=D0=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- static/js/asr_settings.js | 99 ++++++++++++++++++++++++++++++ templates/user/history_detail.html | 1 + 2 files changed, 100 insertions(+) create mode 100644 static/js/asr_settings.js diff --git a/static/js/asr_settings.js b/static/js/asr_settings.js new file mode 100644 index 0000000..33fbd6a --- /dev/null +++ b/static/js/asr_settings.js @@ -0,0 +1,99 @@ +(function() { + 'use strict'; + + const STORAGE_KEY = 'asr_expert_settings_v1'; + const DEFAULTS = { + url: { + expert: false, + keep_raw: true, + do_echo_clearing: false, + do_dialogue: true, + do_punctuation: true, + do_diarization: false, + make_mono: true, + diar_vad_sensity: 3, + do_auto_speech_speed_correction: false, + speech_speed_correction_multiplier: 1.0, + use_batch: false, + batch_size: 8, + fast_speech: false, + split_phrases: true, + }, + file: { + expert: false, + keep_raw: true, + do_echo_clearing: false, + do_dialogue: true, + do_punctuation: true, + do_diarization: false, + make_mono: true, + diar_vad_sensity: 3, + do_auto_speech_speed_correction: false, + speech_speed_correction_multiplier: 1.0, + use_batch: false, + batch_size: 8, + fast_speech: false, + split_phrases: true, + }, + ws: { + expert: false, + do_dialogue: true, + do_punctuation: true, + use_base64: true, + wait_null_answers: true, + audio_transport: 'json_base64', + sample_rate: 16000, + audio_format: 'pcm16', + channel_name: '', + fast_speech: false, + split_phrases: true, + } + }; + + function _load() { + try { + const raw = localStorage.getItem(STORAGE_KEY); + return raw ? JSON.parse(raw) : {}; + } catch (e) { + return {}; + } + } + + function _save(data) { + localStorage.setItem(STORAGE_KEY, JSON.stringify(data)); + } + + function getFor(mode) { + const stored = _load(); + const defs = DEFAULTS[mode] || {}; + return { ...defs, ...(stored[mode] || {}) }; + } + + function setFor(mode, values) { + const stored = _load(); + stored[mode] = { ...(stored[mode] || {}), ...values }; + _save(stored); + } + + function reset(mode) { + const stored = _load(); + if (mode) { + delete stored[mode]; + } else { + Object.keys(stored).forEach(k => delete stored[k]); + } + _save(stored); + } + + function getDefaults(mode) { + return DEFAULTS[mode] || {}; + } + + function isDirty(mode) { + const current = getFor(mode); + const defs = getDefaults(mode); + return Object.keys(defs).some(key => current[key] !== defs[key]); + } + + window.ASRSettings = { getFor, setFor, reset, getDefaults, isDirty }; +})(); diff --git a/templates/user/history_detail.html b/templates/user/history_detail.html index 792950e..e68e240 100644 --- a/templates/user/history_detail.html +++ b/templates/user/history_detail.html @@ -126,6 +126,7 @@

Результат

window.location.href = '/history'; } } + `) в `dashboard.html` и `history_detail.html`, из-за которых не загружались виджеты и детали сессии. | `templates/user/dashboard.html`, `templates/user/history_detail.html` | Страницы загружаются без ошибок в консоли. | +| 6.7.доп.3 | **Вкладки ASR.** Убрана линия под кнопками вкладок (URL/File/WS). Активная вкладка подсвечивается цветом (`rgba(59, 130, 246, 0.2)` + рамка). | `templates/user/asr.html` | Визуально понятно, какая вкладка активна. | +| 6.7.доп.4 | **Кликабельные строки сессий.** Строки таблиц «Последние сессии» на Dashboard и в мини-истории ASR стали кликабельными — клик ведёт на `/history/{id}`. Убрана кнопка «Открыть». | `templates/user/dashboard.html`, `static/js/asr_page.js` | Навигация работает через `window.location.href`. | +| 6.7.доп.5 | **Кнопка «Назад» через history.back().** На странице деталей сессии кнопка «Назад» использует `history.back()` вместо фиксированной ссылки `/history`, что позволяет возвращаться на любую предыдущую страницу (dashboard, asr, history). | `templates/user/history_detail.html` | Корректный переход назад. | +| 6.7.доп.6 | **Автоматические фильтры истории.** Фильтры на странице `/history` применяются автоматически при изменении любого поля (`onchange`/`oninput` + debounce 300мс). Кнопка «Применить» удалена как избыточная. | `templates/user/history.html` | UX фильтрации стал мгновенным. | +| 6.7.доп.7 | **Клиентская фильтрация истории.** Добавлена fallback-фильтрация на фронте (по типу, статусу, дате, поиску по ID) на случай, если бэкенд игнорирует query-параметры. Добавлен `AbortController` для отмены устаревших запросов. | `templates/user/history.html` | Фильтры работают независимо от бэкенда. | +| 6.7.доп.8 | **Перегруппировка переключателей ASR.** Переключатель «Диаризация» перенесён из блока «Пост-обработка» в блок «Диаризация». Создан новый блок «Пред-обработка» с переключателем «Объединить каналы (mono)». | `templates/user/asr.html` | Логическая группировка параметров. | +| 6.7.доп.9 | **Имя файла в мини-истории ASR.** В таблице последних сессий на странице ASR добавлена колонка с именем файла (`file_name`) или последним сегментом URL. | `static/js/asr_page.js` | Пользователь видит, какой файл распознавался. | +| 6.7.доп.10 | **Переключение формата результата (DIALOG/TEXT/RAW).** Реализованы кнопки DIALOG, TEXT, RAW для переключения формата отображения результата распознавания на всех вкладках (URL, File, WS) и на странице деталей сессии. По умолчанию: DIALOG (если есть диаризация/фразы), иначе TEXT. RAW — подсвеченный JSON. | `static/js/asr_page.js`, `templates/user/asr.html`, `templates/user/history_detail.html` | Пользователь выбирает удобный формат просмотра. | +| 6.7.доп.11 | **Единая кнопка «Скачать».** Кнопки «Скачать JSON» и «Скачать TXT» объединены в одну «Скачать». Формат файла определяется текущим режимом отображения: RAW → `.json`, TEXT/DIALOG → `.txt`. | `static/js/asr_page.js`, `templates/user/asr.html`, `templates/user/history_detail.html` | Скачивается то, что видит пользователь. | +| 6.7.доп.12 | **Модуль настроек ASR (`asr_settings.js`).** Создан `static/js/asr_settings.js` — модуль хранения настроек режима Эксперт в `localStorage`. Исправлен баг `ASRSettings is not defined` на странице ASR. | `static/js/asr_settings.js` | Настройки эксперта сохраняются между сессиями. | + +--- + +### Шаг 6.8. Запись аудио с микрофона (WebRTC MediaRecorder) +**Цель:** Дать пользователю возможность записывать аудио прямо в браузере и отправлять его на распознавание без загрузки файла. + +| Файл | Действие | +|------|----------| +| `templates/user/asr.html` | **Доработать.** Добавлена четвёртая вкладка «🎙 Микрофон». UI состоит из трёх состояний: (1) idle — большая круглая кнопка «🎙» для начала записи; (2) recording — пульсирующая красная точка + таймер `MM:SS` + кнопка «⏹ Стоп»; (3) preview — `