diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..cfe055c --- /dev/null +++ b/.dockerignore @@ -0,0 +1,16 @@ +venv/ +.venv/ +models/ +__pycache__/ +*.py[cod] +.git/ +logs/ +trash/ +bench/ +*.log +*.db +asr_local.db +.idea/ +.dockerignore +local_asr/to_asr/* +local_asr/after_asr/* diff --git a/.gitignore b/.gitignore index 6d883c2..40e4ccf 100644 --- a/.gitignore +++ b/.gitignore @@ -4,7 +4,7 @@ *.ogg *.onnx *.m4a -*.md +*.history.md /.venv/ /venv/ /models/vosk-model-ru/ @@ -17,3 +17,16 @@ /models/sherpa-onnx-whisper-small/ /trash/test_data.py /.idea/ +.aider* +/.continue/ +.env + +# runtime / cache / dev +__pycache__/ +*.py[cod] +*.db +asr_local.db +/logs/ +/bench/ +/models/hub/ +/models/sbert_punc_case_ru_onnx/ \ No newline at end of file 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/Diarisation/__init__.py b/Diarisation/__init__.py index f879a5c..bab5650 100644 --- a/Diarisation/__init__.py +++ b/Diarisation/__init__.py @@ -1,27 +1,29 @@ -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 +import logging +from starlette.requests import HTTPConnection +logger = logging.getLogger(__name__) -if config.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"Модель для диаризации не найдена. Предпринимаются попытки скачать {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" - try: # Предпринимаем попытки скачать и положить файл в нужное место. - # 1. Получаем JSON с данными файлов + try: 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} в списке возможных для загрузки нет.") - # Получаем список всех ONNX-моделей + logger.error(f"Модели с именем {settings.DIAR_MODEL_NAME} в списке возможных для загрузки нет.") 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"Будет использован имеющийся файл {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"Модель для Диаризации отсутствует. Диаризация выключена.") 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 - 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=config.DIAR_GPU_BATCH_SIZE, - cpu_workers=config.CPU_WORKERS, - use_gpu=config.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 de7f6df..2915c9e 100644 --- a/Diarisation/diarazer.py +++ b/Diarisation/diarazer.py @@ -1,18 +1,16 @@ import datetime - -import config +from pydub import AudioSegment 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 config.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, @@ -193,6 +191,38 @@ def match_asr_with_diarization(asr_data, diarization_data, min_overlap_ratio=0.5 return dict(result) +async def do_diarizing_v1( + audio_segment: AudioSegment, + asr_raw_data, + diarizer, + num_speakers:int = -1, + filter_cutoff:int = 50, + filter_order:int = 10, + diar_vad_sensity: int = 3 + ): + # Предобработка аудио. + # ВАЖНЫЙ момент. Мы диаризируем только последний канал. + audio_frames = await load_and_preprocess_audio(audio_segment.split_to_mono()[-1]) + + logger.debug("Старт процесса диаризации.") + st=datetime.datetime.now() + # Непосредственно получение временных меток речи + diar_result = await diarizer.diarize_and_merge( + audio_frames, + num_speakers=num_speakers, + filter_cutoff=filter_cutoff, + filter_order=filter_order, + vad_sensity=diar_vad_sensity + ) + + for r in diar_result: + logger.debug(f"Спикер {r['speaker']}: {r['start']:.2f} - {r['end']:.2f} сек") + logger.debug(f"Диаризация завершена за {(datetime.datetime.now()-st).total_seconds()} секунд") + + # Построение структуры аналогично raw_data для дальнейшего построения диалога и вывод результата + return match_asr_with_diarization(asr_raw_data, diar_result) + + def group_words(words): if not words: return [] @@ -219,4 +249,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 a642d61..b4d71fa 100644 --- a/Diarisation/do_diarize.py +++ b/Diarisation/do_diarize.py @@ -13,9 +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: diff --git a/Docker/Dockerfile.local b/Docker/Dockerfile.local new file mode 100644 index 0000000..ba7a821 --- /dev/null +++ b/Docker/Dockerfile.local @@ -0,0 +1,47 @@ +# Локальный ТЕСТОВЫЙ образ новой архитектуры (api/v1 + services): собирается из рабочего +# дерева (COPY), модели монтируются с хоста. Прод-образ - Docker/Dockerfile. +# +# Сборка из КОРНЯ репозитория: +# CPU: docker build -f Docker/Dockerfile.local -t asr-local . +# GPU: docker build -f Docker/Dockerfile.local --build-arg PROVIDER=CUDA -t asr-local . +FROM python:3.12-slim-bullseye + +# build-essential/cmake/libboost/libeigen нужны для сборки kenlm (зависимость T-one) +RUN apt-get update && apt-get install -y \ + git git-lfs ffmpeg curl ca-certificates jq \ + build-essential cmake libboost-all-dev libeigen3-dev \ + && rm -rf /var/lib/apt/lists/* +RUN git lfs install + +# Провайдер: CPU (по умолч.) | CUDA. Для CUDA меняем onnxruntime -> onnxruntime-gpu + cupy. +ARG PROVIDER=CPU +COPY requirements.txt /tmp/req.txt +RUN if [ "$PROVIDER" = "CUDA" ] || [ "$PROVIDER" = "TENSORRT" ]; then \ + sed -i 's#^onnxruntime\s*$#onnxruntime-gpu[cuda,cudnn]==1.23.2#' /tmp/req.txt && \ + echo "cupy-cuda12x==13.5.1" >> /tmp/req.txt ; \ + fi +RUN pip install --no-cache-dir -r /tmp/req.txt + +# T-one (потоковый путь /api/v1/asr/ws-stream): --no-deps, затем зависимости (kenlm соберётся) +RUN pip install --no-cache-dir --no-deps "git+https://github.com/voicekit-team/T-one.git" && \ + pip install --no-cache-dir pyctcdecode kenlm miniaudio + +# Для CUDA фиксируем cuDNN как в проде +RUN if [ "$PROVIDER" = "CUDA" ] || [ "$PROVIDER" = "TENSORRT" ]; then \ + pip install --no-cache-dir --force-reinstall nvidia-cudnn-cu12==9.6.0.74 ; \ + fi + +# Код приложения - из локального рабочего дерева (.dockerignore исключает venv/models/.git) +COPY . /ASR_FastAPI_WS_RU/ + +# Каталоги, обязательные на старте (логи; иначе TimedRotatingFileHandler падает) +RUN mkdir -p /ASR_FastAPI_WS_RU/logs /ASR_FastAPI_WS_RU/local_asr/to_asr /ASR_FastAPI_WS_RU/local_asr/after_asr + +ENV PYTHONPATH=/ASR_FastAPI_WS_RU \ + HF_HOME=/ASR_FastAPI_WS_RU/models \ + IS_PROD=1 \ + LOGGING_LEVEL=INFO + +WORKDIR /ASR_FastAPI_WS_RU +EXPOSE 49153 +CMD ["python3", "main.py"] diff --git a/Punctuation/__init__.py b/Punctuation/__init__.py index 06c912e..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 -import config - -try: - sbertpunc = SbertPuncCaseOnnx(paths.get("punctuation_model_path"), use_gpu = config.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 92e3923..5819266 100644 --- a/Recognizer/__init__.py +++ b/Recognizer/__init__.py @@ -1,14 +1,16 @@ +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 -import config +from config import settings import onnxruntime as ort import onnx_asr from onnx_asr.loader import PreprocessorRuntimeConfig, OnnxSessionOptions +logger = logging.getLogger(__name__) TENSORRT_providers = ["TensorrtExecutionProvider", "CUDAExecutionProvider", "CPUExecutionProvider"] CUDA_providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] @@ -18,14 +20,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 +106,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("Ошибка при прогреве модели. Сервис работать не будет. Возможно, модель не поддерживает выбранный провайдер.") @@ -122,4 +124,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/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 635e315..30477ef 100644 --- a/Recognizer/engine/file_recognition.py +++ b/Recognizer/engine/file_recognition.py @@ -1,30 +1,30 @@ import time - from pydub import AudioSegment -import config +from config import settings import asyncio -import uuid -from utils.pre_start_init import ( - posted_and_downloaded_audio, - audio_buffer, - audio_overlap, - audio_to_asr, - audio_duration, -) -from utils.do_logging import logger -from utils.chunk_doing import find_last_speech_position +import logging + +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 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 -from Diarisation.diarazer import do_diarizing -from threading import Lock +from Diarisation.diarazer import do_diarizing, do_diarizing_v1 +from utils.pre_start_init import posted_and_downloaded_audio + +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 + legacy_mode = False + if session is None and tmp_path is not None and params is not None: + session = FileRecognitionSession(params=params) + session.tmp_path = tmp_path + legacy_mode = True + elif session is None: + raise TypeError("process_file() требует либо 'session', либо оба 'tmp_path' и 'params'") -def process_file(tmp_path, params): process_file_start = time.perf_counter() res = False diarized = False @@ -37,15 +37,21 @@ def process_file(tmp_path, params): "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) + 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(file_source).set_channels(1) + else: + audio_input = AudioSegment.from_file(file_source) + + if legacy_mode: + posted_and_downloaded_audio[session.post_id] = audio_input except Exception as e: error_description += f"Error loading audio file: {e}" logger.error(error_description) @@ -55,10 +61,9 @@ def process_file(tmp_path, params): # Проверка длины переданного на распознавание аудио 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=config.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) @@ -67,21 +72,14 @@ def process_file(tmp_path, params): 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 != config.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) - print(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) @@ -90,114 +88,105 @@ def process_file(tmp_path, params): 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=config.BASE_SAMPLE_RATE) - audio_overlap[post_id] = AudioSegment.silent(1, frame_rate=config.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[::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) - 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) # --> list + list_asr_result_wo_conf = asyncio.run(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_asr, - can_slow_down=True, - multiplier=params.speech_speed_correction_multiplier)) + 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) + ) 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}") 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}") 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 # Проверяем возможность диаризации. Если здесь стерео-канал, то диаризацию выключаем. - 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("При запрошенной диаризации аудио имеет более одного аудио-канала. Диаризация будет выключена.") + logger.warning("При запрошенной диаризации аудио имеет более одного аудио-канала. Диаризация будет выключена.") params.do_diarization = False 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 - )) + if legacy_mode: + result["diarized_data"] = asyncio.run(do_diarizing( + file_id=str(post_id), + asr_raw_data=session.collected_asr_res, + diar_vad_sensity=params.diar_vad_sensity, + diarizer=diarizer, + )) + else: + result["diarized_data"] = asyncio.run(do_diarizing_v1( + audio_segment=audio_input, + asr_raw_data=session.collected_asr_res, + diar_vad_sensity=params.diar_vad_sensity, + diarizer=diarizer, + )) except Exception as e: logger.error(f"do_diarizing - {e}") error_description = f"do_diarizing - {e}" @@ -206,27 +195,28 @@ def process_file(tmp_path, params): 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 - ) - ) + 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/Recognizer/engine/sentensizer.py b/Recognizer/engine/sentensizer.py index ae9683c..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 -import config -from Punctuation import sbertpunc +from config import settings +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, @@ -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 @@ -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 13d4dd5..5b57f61 100644 --- a/Recognizer/engine/stream_recognition.py +++ b/Recognizer/engine/stream_recognition.py @@ -1,15 +1,16 @@ +import asyncio import time import numpy as np -import config -from Recognizer import recognizer +from config import settings 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 @@ -41,23 +42,29 @@ def calc_speed(data): return speech_speed -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) +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) # Перевод в семплы для распознавания. 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 +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, - ) -> tuple: + recognizer = None) -> tuple: """ Распознавание чанка с возможностью контроля быстрой речи. + :param recognizer: Обязательно получить класс модели распознавания. :param multiplier: Float :param can_slow_down: Boolean :param audio_data: Аудиоданные в формате Audiosegment (puDub). @@ -71,22 +78,23 @@ 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) 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 -def simple_recognise_batch(list_audio_data: list, batch_size: int = 8) -> 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() @@ -97,7 +105,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) # Делаем препроцессинг на все данные сразу @@ -134,9 +142,12 @@ def simple_recognise_batch(list_audio_data: list, batch_size: int = 8) -> list: 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) @@ -150,4 +161,4 @@ def simple_recognise_batch(list_audio_data: list, batch_size: int = 8) -> list: "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/Recognizer/tone_engine.py b/Recognizer/tone_engine.py new file mode 100644 index 0000000..46e227c --- /dev/null +++ b/Recognizer/tone_engine.py @@ -0,0 +1,66 @@ +# -*- coding: utf-8 -*- +""" +Ленивый синглтон потокового движка T-one (tone.StreamingCTCPipeline). + +Грузится один раз при первом обращении (а не при импорте), чтобы не тянуть модель, +пока потоковый эндпоинт не используется, и не вмешиваться в загрузку офлайн-модели +GigaAM. Модель и KenLM качаются с HF (t-tech/T-one) в каталог HF_HOME. + +Провайдер акустической модели: + - по умолчанию CPU (так задумано библиотекой; для чанков по 300 мс это обычно оптимально); + - при settings.STREAM_WITH_GPU=1 и GPU-провайдере собираем pipeline вручную с + CUDAExecutionProvider (даёт возможность померить, выгоден ли GPU на коротких чанках). +""" + +import logging + +from config import settings + +logger = logging.getLogger(__name__) + +_pipeline = None + +CUDA_PROVIDERS = ["CUDAExecutionProvider", "CPUExecutionProvider"] + + +def _build_gpu_pipeline(): + """Собирает StreamingCTCPipeline с CUDA-сессией акустической модели.""" + import onnxruntime as ort + from tone import StreamingCTCPipeline, DecoderType + from tone.onnx_wrapper import StreamingCTCModel + from tone.logprob_splitter import StreamingLogprobSplitter + from tone.decoder import GreedyCTCDecoder, BeamSearchCTCDecoder + + model_path = StreamingCTCModel.download_from_hugging_face() + sess = ort.InferenceSession(model_path, providers=CUDA_PROVIDERS) + model = StreamingCTCModel(sess) + splitter = StreamingLogprobSplitter() + if str(settings.TONE_DECODER).lower() == "greedy": + decoder = GreedyCTCDecoder() + else: + decoder = BeamSearchCTCDecoder.from_hugging_face() + logger.info("Использован провайдер %s для T-one", sess.get_providers()[0]) + return StreamingCTCPipeline(model, splitter, decoder) + + +def get_tone_pipeline(): + """Возвращает singleton StreamingCTCPipeline, загружая его при первом вызове.""" + global _pipeline + if _pipeline is None: + from tone import StreamingCTCPipeline, DecoderType + + decoder_type = (DecoderType.GREEDY + if str(settings.TONE_DECODER).lower() == "greedy" + else DecoderType.BEAM_SEARCH) + logger.info("Загрузка потоковой модели T-one (decoder=%s, gpu=%s)...", + decoder_type.value, settings.STREAM_WITH_GPU) + if settings.STREAM_WITH_GPU: + try: + _pipeline = _build_gpu_pipeline() + except Exception as exc: + logger.error("Не удалось поднять T-one на GPU (%s), откат на CPU", exc) + _pipeline = StreamingCTCPipeline.from_hugging_face(decoder_type=decoder_type) + else: + _pipeline = StreamingCTCPipeline.from_hugging_face(decoder_type=decoder_type) + logger.info("Потоковая модель T-one загружена.") + return _pipeline diff --git a/VoiceActivityDetector/__init__.py b/VoiceActivityDetector/__init__.py index 751633b..dfe071b 100644 --- a/VoiceActivityDetector/__init__.py +++ b/VoiceActivityDetector/__init__.py @@ -1,8 +1,8 @@ -import config +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 отсутствует. Предпринимаем попытку скачать её.") @@ -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/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/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/deps.py b/api/deps.py new file mode 100644 index 0000000..c01908b --- /dev/null +++ b/api/deps.py @@ -0,0 +1,99 @@ +from datetime import datetime, timezone +from typing import Optional, Any + +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/api/legacy/__init__.py b/api/legacy/__init__.py new file mode 100644 index 0000000..ba1557b --- /dev/null +++ b/api/legacy/__init__.py @@ -0,0 +1,19 @@ +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 +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 + +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) +router.include_router(is_alive_router) +router.include_router(post_by_url_router) +router.include_router(post_by_file_router) diff --git a/routes/demo_page.py b/api/legacy/demo_page.py similarity index 52% rename from routes/demo_page.py rename to api/legacy/demo_page.py index fa01f14..0dcf201 100644 --- a/routes/demo_page.py +++ b/api/legacy/demo_page.py @@ -1,20 +1,20 @@ -from utils.pre_start_init import app -from fastapi import WebSocket, WebSocketException, Request -from utils.do_logging import logger -from fastapi.staticfiles import StaticFiles +import logging +from fastapi import APIRouter, Request from fastapi.responses import HTMLResponse from fastapi.templating import Jinja2Templates +logger = logging.getLogger(__name__) - -# 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): + 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 new file mode 100644 index 0000000..a2480dc --- /dev/null +++ b/api/legacy/is_alive.py @@ -0,0 +1,62 @@ +from fastapi import APIRouter +import logging +import pynvml + +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: + 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)}, None, None, None + return None, free_mb, gpu_load, temperature + + +@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') + # Legacy: глобальный audio_to_asr больше не используется в новой архитектуре + tasks_in_work = 0 + + 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 {"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/api/legacy/post_by_file_FORM.py similarity index 59% rename from routes/post_by_file_FORM.py rename to api/legacy/post_by_file_FORM.py index 5d8dce9..79f8bdf 100644 --- a/routes/post_by_file_FORM.py +++ b/api/legacy/post_by_file_FORM.py @@ -1,17 +1,18 @@ from io import BytesIO import asyncio - -import config -from utils.pre_start_init import app -from utils.do_logging import logger -from models.fast_api_models import PostFileRequest +import logging +from config import settings +from fastapi import APIRouter, Depends, File, Form, UploadFile +from models.fast_api_models import PostFileRequest, BaseResponse +from Recognizer import get_recognizer, Recognizer from Recognizer.engine.file_recognition import process_file -from fastapi import Depends, File, Form, UploadFile -from threading import Lock +from Punctuation import get_punctuator, SbertPuncCaseOnnx +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer -# Глобальный лок для потокобезопасности -audio_lock = Lock() +logger = logging.getLogger(__name__) +router = APIRouter() # Функция для извлечения параметров из FormData def get_file_request( @@ -21,8 +22,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"), @@ -42,21 +43,17 @@ def get_file_request( ) -@app.post("/post_file") -async def async_receive_file( +@router.post("/post_file", response_model=BaseResponse) +async def async_receive_file_legacy( 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(), - } - + recognizer: Recognizer = Depends(get_recognizer), + 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()) @@ -64,19 +61,34 @@ 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, + tmp_path=buffer, + params=params, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer) + 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/api/legacy/post_by_url.py b/api/legacy/post_by_url.py new file mode 100644 index 0000000..1fe4440 --- /dev/null +++ b/api/legacy/post_by_url.py @@ -0,0 +1,83 @@ +import uuid +import asyncio +from fastapi import APIRouter, Depends +from utils.get_audio_file import getting_audiofile, open_default_audiofile +from models.fast_api_models import SyncASRRequest, BaseResponse + +from Recognizer import get_recognizer, Recognizer +from Punctuation import get_punctuator, SbertPuncCaseOnnx +from Recognizer.engine.file_recognition import process_file + +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer + +import logging +logger = logging.getLogger(__name__) +router = APIRouter() + +@router.post("/post_one_step_req", response_model=BaseResponse) +async def post(params: SyncASRRequest, + recognizer: Recognizer = Depends(get_recognizer), + punctuator: SbertPuncCaseOnnx = Depends(get_punctuator), + diarizer: Diarizer = Depends(get_diarizer) +) -> BaseResponse: + """ + На вход ждёт str(HttpUrl) - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n + Если на вход передаётся не моно, то ответ будет в несколько элементов списка для каждого канала.\n + + :param: do_dialogue: - true, если нужно разбить речь на диалог\n + :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 + """ + + logger.warning( + "Legacy endpoint /post_one_step_req is deprecated. Use /api/v1/asr/url instead.", + ) + + # Получаем файл + post_id = uuid.uuid4() + if params.AudioFileUrl: + res, error_description, buffer = await getting_audiofile(params.AudioFileUrl, post_id) + else: + res, error_description, buffer = await open_default_audiofile(post_id) + + if not res: + logger.error(f'Ошибка получения файла - {error_description}, ссылка на файл - {params.AudioFileUrl}') + return BaseResponse( + success=False, + error_description=error_description, + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) + + try: + # Запускаем обработку в потоке + result_dict = await asyncio.to_thread(process_file, + tmp_path=buffer, + params=params, + recognizer=recognizer, + punctuator=punctuator, + diarizer=diarizer) + + result = BaseResponse(**result_dict) + except Exception as e: + error_description = f"Ошибка обработки в process_file - {e}" + logger.error(error_description) + return BaseResponse( + success=False, + error_description=str(error_description), + raw_data={}, + sentenced_data={}, + diarized_data={}, + ) + + return result diff --git a/api/legacy/root.py b/api/legacy/root.py new file mode 100644 index 0000000..76e7bd7 --- /dev/null +++ b/api/legacy/root.py @@ -0,0 +1,29 @@ +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("/") +async def root(): + """ + Корневой эндпоинт API v1. + + 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", + "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/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..7377c86 --- /dev/null +++ b/api/v1/api.py @@ -0,0 +1,23 @@ +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.asr_ws_tone import router as asr_ws_tone_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") + +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(asr_ws_tone_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/__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/api/v1/endpoints/admin.py b/api/v1/endpoints/admin.py new file mode 100644 index 0000000..ed2042b --- /dev/null +++ b/api/v1/endpoints/admin.py @@ -0,0 +1,668 @@ +"""FastAPI-роутер админ-панели.""" + +from collections import defaultdict +from datetime import datetime, timedelta +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 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, + 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( + 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]) +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, + 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, + 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, + 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, + "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.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 = 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), + 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(), + user_id: Optional[str] = None, + current_user: User = Depends(require_admin), + db: AsyncSession = Depends(get_db_session), +): + """Список подписок.""" + 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, + user_id=s.user_id, + plan_id=s.plan_id, + plan_name=None, + status=s.status, + 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, + 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, + 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..592aa3b --- /dev/null +++ b/api/v1/endpoints/admin_ws.py @@ -0,0 +1,72 @@ +"""WebSocket endpoint для real-time метрик админ-панели.""" + +import asyncio +import logging + +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 +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", "") + 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 + + 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), + ) + 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"}} + ) + 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 new file mode 100644 index 0000000..6eeb980 --- /dev/null +++ b/api/v1/endpoints/asr_file.py @@ -0,0 +1,116 @@ +import asyncio +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 +from Recognizer import get_recognizer, Recognizer +from Punctuation import get_punctuator, SbertPuncCaseOnnx + +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer + +logger = logging.getLogger(__name__) +router = APIRouter(prefix="/asr", tags=["ASR"]) + + +# Функция для извлечения параметров из 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("/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), + 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}") + result_dict = await asyncio.to_thread( + process_file, + session=session, + recognizer=recognizer, + 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'), + 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() + session.cleanup() + await session.reset() diff --git a/api/v1/endpoints/asr_url.py b/api/v1/endpoints/asr_url.py new file mode 100644 index 0000000..c2074df --- /dev/null +++ b/api/v1/endpoints/asr_url.py @@ -0,0 +1,110 @@ +import uuid +import asyncio +import logging +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 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 +from Diarisation import get_diarizer +from Diarisation.do_diarize import Diarizer + +logger = logging.getLogger(__name__) + +router = APIRouter(prefix="/asr", tags=["ASR"]) + + +@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), + 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) + 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 file_buffer is None: + return V1BaseResponse( + success=False, + error_description="Получен пустой буфер файла", + data=ASRData() + ) + session.file_buffer = file_buffer + session.file_buffer.seek(0) + + result_dict = await asyncio.to_thread( + process_file, + session=session, + recognizer=recognizer, + 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'), + 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: + session.cleanup() + await session.reset() diff --git a/api/v1/endpoints/asr_ws.py b/api/v1/endpoints/asr_ws.py new file mode 100644 index 0000000..f929f35 --- /dev/null +++ b/api/v1/endpoints/asr_ws.py @@ -0,0 +1,280 @@ +""" +WebSocket-роут /api/v1/asr/ws +Использует ConnectionManager, AudioSession, MessageRouter, asr_pipeline. +Сохраняет обратную совместимость протокола (config, audio, eof/eos). +""" + +import asyncio +import base64 +import logging +import uuid +from contextlib import asynccontextmanager + +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_protocol import normalize_to_ws_message +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 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__) + + +@asynccontextmanager +async def audio_session_lifecycle(client_id: str): + """ + Контекстный менеджер жизненного цикла AudioSession. + + Гарантирует очистку AudioSegment-буферов. + """ + session = AudioSession(client_id=client_id) + try: + yield session + finally: + await session.reset() + + +@router.websocket("/ws") +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). + + Протокол: + 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 + + # 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) + msg_router.register_handler(WSMessageType.ping, handle_ping) + msg_router.register_handler(WSMessageType.status_request, handle_status_request) + + try: + while True: + # 4. Получение сообщения с idle timeout + try: + 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( + client_id, + WSErrorMessage( + code="idle_timeout", + message=f"No messages for {settings.WS_IDLE_TIMEOUT_SEC}s", + is_fatal=False, + ), + ) + break + + # Обработка disconnect от клиента + if message.get("type") == "websocket.disconnect": + logger.info("Client %s disconnected (code=%s)", client_id, message.get("code")) + break + + # Определяем тип содержимого: 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 + ) + continue + elif message.get("text"): + # Автодетект протокола: понимает и новый ({type:...}), и legacy ({config}/{eof}) + msg = normalize_to_ws_message(message["text"]) + if msg is None: + logger.warning("Unrecognized WS text message from %s", client_id) + continue + else: + logger.warning("Unknown WS message format for %s: %s", client_id, message) + continue + + # 5. Маршрутизация служебных сообщений + if msg.type in ( + WSMessageType.config, + WSMessageType.ping, + WSMessageType.status_request, + ): + 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.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) + if chunk_bytes: + await process_audio_stream_chunk( + session, chunk_bytes, recognizer, punctuator, manager, metrics + ) + + # 7. Обработка конца потока (eos/eof) + if isinstance(msg, WSEosMessage): + await process_final_audio(session, recognizer, punctuator, manager, metrics) + break + + except WebSocketDisconnect: + 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, + WSErrorMessage( + code="internal_error", + message=str(exc), + is_fatal=True, + ), + ) + except Exception: + 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()) + except Exception as exc: + logger.debug("Failed to save session state: %s", exc) + await manager.disconnect(client_id) diff --git a/api/v1/endpoints/asr_ws_tone.py b/api/v1/endpoints/asr_ws_tone.py new file mode 100644 index 0000000..77d5ddb --- /dev/null +++ b/api/v1/endpoints/asr_ws_tone.py @@ -0,0 +1,146 @@ +# -*- coding: utf-8 -*- +""" +WebSocket-роут /api/v1/asr/ws-stream — НАСТОЯЩИЙ потоковый риалтайм на нативном T-one. + +В отличие от /api/v1/asr/ws (офлайн псевдо-стрим GigaAM: копит до MAX_OVERLAP_DURATION и +распознаёт целым куском), здесь аудио скармливается модели кадрами по 300 мс с сохранением +состояния, а фразы отдаются по мере их завершения детектором границ T-one (~0.3-1 c). + +Особенности: + - Автодетект протокола (services/ws_protocol): понимает и legacy ({config}/{eof}/raw bytes), + и новый ({type:config/audio_chunk/eos/ping}). Ответы в формате WSResultMessage — + его поля (silence/data/last_message) совместимы с legacy asterisk-socket-server. + - Потоковый ресемплинг источник->8 кГц (T-one фиксирован на 8 кГц). + - Времена фраз накопительны от первого пакета (по объёму поданного аудио) — корректны + для поканального мёржа на стороне клиента. + - БД не используется (в отличие от asr_ws.py) — эндпоинт независим от наличия таблиц. +""" + +import logging +import uuid + +from fastapi import APIRouter, WebSocket + +from config import settings +from models.ws_models import ( + WSResultMessage, + WSRecognitionData, + WSPongMessage, + WSMessageType, +) +from services.ws_manager import ConnectionManager +from services.ws_protocol import detect +from utils.tone_stream import take_frames, flush_tail, phrase_to_data, StreamResampler +from Recognizer.tone_engine import get_tone_pipeline + +router = APIRouter(prefix="/asr", tags=["ASR"]) +logger = logging.getLogger(__name__) + + +def _result_message(phrase, channel_name: str, last: bool = False) -> WSResultMessage: + return WSResultMessage( + type=WSMessageType.final_result if last else WSMessageType.partial_result, + channel_name=channel_name, + silence=False, + data=WSRecognitionData(**phrase_to_data(phrase)), + last_message=last, + ) + + +@router.websocket("/ws-stream") +async def websocket_tone_stream(websocket: WebSocket): + manager: ConnectionManager = websocket.app.state.ws_manager + client_id = str(uuid.uuid4()) + + if not await manager.connect(websocket, client_id): + return # лимит соединений исчерпан, ConnectionManager уже закрыл сокет + + # Предзагруженный в lifespan движок (готов сразу после старта); фолбэк - ленивая загрузка + pipeline = getattr(websocket.app.state, "tone_pipeline", None) or get_tone_pipeline() + state = None + buf = bytearray() + channel_name = "Null" + resampler = StreamResampler(settings.TONE_SAMPLE_RATE) # проходной, пока не пришёл config + + logger.debug("[tone] new stream %s", client_id) + + try: + while True: + try: + message = await websocket.receive() + except Exception as exc: + logger.debug("[tone] receive error %s: %s", client_id, exc) + break + + evt = detect(message) + + if evt.kind == "disconnect": + logger.info("[tone] disconnect %s (%s)", channel_name, client_id) + break + + if evt.kind == "config": + channel_name = evt.channel_name or "Null" + sr = evt.sample_rate or settings.TONE_SAMPLE_RATE + resampler = StreamResampler(sr) + if sr != settings.TONE_SAMPLE_RATE: + logger.info("[tone] %s: sample_rate=%s, включён ресемплинг к 8 кГц", channel_name, sr) + logger.info("[tone] config received for channel %s", channel_name) + continue + + if evt.kind == "ping": + await manager.send_message(client_id, WSPongMessage()) + continue + + if evt.kind == "audio": + # Время T-one считается по объёму поданного аудио (от первого пакета), + # ресемплинг сохраняет длительность — метки остаются корректными. + buf.extend(resampler.process(evt.audio or b"")) + try: + for samples in take_frames(buf): + phrases, state = pipeline.forward(samples, state) + for phrase in phrases: + if phrase.text: + await manager.send_message(client_id, _result_message(phrase, channel_name)) + except Exception as exc: + logger.error("[tone] recognize error %s (%s): %s", channel_name, client_id, exc) + continue + + if evt.kind == "eos": + logger.info("[tone] EOS for channel %s", channel_name) + break + + # evt.kind == "ignore" — молча пропускаем + + # --- финализация: дослать хвост ресемплера и буфера, закрыть фразы --- + last_phrase = None + try: + buf.extend(resampler.process(b"", last=True)) + tail = flush_tail(buf) + final_phrases = [] + if tail is not None: + phrases, state = pipeline.forward(tail, state, is_last=True) + final_phrases.extend(phrases) + fin_phrases, state = pipeline.finalize(state) + final_phrases.extend(fin_phrases) + final_phrases = [p for p in final_phrases if p.text] + # все, кроме последней, шлём обычными; последнюю пометим last_message + for p in final_phrases[:-1]: + await manager.send_message(client_id, _result_message(p, channel_name)) + if final_phrases: + last_phrase = final_phrases[-1] + except Exception as exc: + logger.error("[tone] finalize error %s (%s): %s", channel_name, client_id, exc) + + if last_phrase is not None: + await manager.send_message(client_id, _result_message(last_phrase, channel_name, last=True)) + else: + await manager.send_message(client_id, WSResultMessage( + type=WSMessageType.final_result, + channel_name=channel_name, + silence=True, + data=WSRecognitionData(), + last_message=True, + )) + finally: + await manager.disconnect(client_id) + logger.info("[tone] closed %s (%s)", channel_name, client_id) diff --git a/api/v1/endpoints/auth.py b/api/v1/endpoints/auth.py new file mode 100644 index 0000000..f1f1da1 --- /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 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 ( + 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/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..601e54a --- /dev/null +++ b/api/v1/endpoints/root.py @@ -0,0 +1,49 @@ +from fastapi import APIRouter, Request +from models.fast_api_models import V1BaseResponse +from services.ws_metrics import SystemMetricsCollector +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" + } + ) + + +@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() + ) 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/config.py b/config.py index cf19b32..a860ce9 100644 --- a/config.py +++ b/config.py @@ -1,93 +1,357 @@ +import logging 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, timedelta +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 + 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" + MODEL_NAME: str = "gigaam-v3-ctc" + # Стрим из астериска отдаёт только 8к + BASE_SAMPLE_RATE: int = 16000 + PROVIDER: str = "CPU" + NUM_THREADS: int = 0 + + # T-one streaming settings (потоковый риалтайм на /api/v1/asr/ws-stream через нативный T-one) + USE_TONE_STREAMING: bool = False # зарезервировано: занять ли T-one ещё и legacy /ws + TONE_DECODER: str = "beam_search" # "beam_search" (с KenLM) | "greedy" + TONE_SAMPLE_RATE: int = 8000 # фиксировано акустической моделью T-one + TONE_CHUNK_SAMPLES: int = 2400 # 300 мс @ 8 кГц - фиксированный размер чанка T-one + # Гонять акустическую модель T-one на GPU (CUDAExecutionProvider). По умолчанию CPU: + # для чанков по 300 мс GPU может быть медленнее из-за оверхеда на копирования/запуск ядра. + # Актуально только при PROVIDER in (CUDA, TENSORRT). + STREAM_WITH_GPU: bool = False + + # 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-started_{date.today()}') + FILEMODE: str = 'a' + # Срок хранения логов в днях + LOG_BACKUP_COUNT: int = 60 + 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 + # 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 + CAN_PUNCTUATE: bool = True + PUNCTUATE_WITH_GPU: bool = False + + # 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 + + # 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 + + # 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', + 'DIAR_WITH_GPU', 'DO_SPEED_SPEECH_CORRECTION', + 'DO_LOCAL_FILE_RECOGNITIONS', 'DELETE_LOCAL_FILE_AFTR_ASR', + 'HUMAN_FORMAT_MD_FILE', 'USE_TONE_STREAMING', 'STREAM_WITH_GPU', + 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): + # HF_HOME должен быть установлен ДО импорта библиотек, использующих HuggingFace Hub + 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 + + # VAD_WITH_GPU актуален только при GPU-провайдерах + if self.VAD_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.VAD_WITH_GPU = False + + # PUNCTUATE_WITH_GPU актуален только при GPU-провайдерах + if self.PUNCTUATE_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.PUNCTUATE_WITH_GPU = False + + # STREAM_WITH_GPU (T-one на GPU) актуален только при GPU-провайдерах + if self.STREAM_WITH_GPU and self.PROVIDER not in ["CUDA", "TENSORRT"]: + self.STREAM_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." + ) + + # Валидация 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 + + +# Единственный экземпляр настроек +settings = Settings() + + + +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 — `/api/v1/asr/ws` (актуальный протокол) + +### Пример конфигурации + +Отправьте JSON с конфигурацией: + +```json +{ + "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) + +```json +{ + "type": "audio_chunk", + "audio_base64": "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAEAfAAABAAgAZGF0YQAAAAA=", + "seq_num": 0 +} +``` + +### Пример передачи аудио (Binary) + +Используется при `audio_transport: "binary"`. Отправляйте WebSocket **binary frame** напрямую (без JSON-обёртки). Сервер читает его через `receive_bytes()`. + +### Пример завершения потока (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 +{ + "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": ["..."] + } +} +``` + +--- + +## 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" + } +} +``` + +### Пример передачи данных (legacy) + +```json +{ + "text": "eof" +} +``` + +### Ответы (legacy) -print(f"Using '{LOGGING_LEVEL}' LOGGING_LEVEL") \ No newline at end of file +```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/__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/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/deps.py b/core/deps.py new file mode 100644 index 0000000..5fb8f07 --- /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 str(user.role) not in (UserRole.admin.value, UserRole.superadmin.value): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Требуются права администратора", + ) + return user + + +def require_superadmin(user: User = Depends(get_current_user)) -> User: + """Проверяет, что пользователь — суперадмин.""" + if str(user.role) != UserRole.superadmin.value: + 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/exception_handlers.py b/core/exception_handlers.py new file mode 100644 index 0000000..bf7b818 --- /dev/null +++ b/core/exception_handlers.py @@ -0,0 +1,60 @@ +import logging +import traceback + +from fastapi import Request +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse, RedirectResponse +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): + 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, + ).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..2e6aa1e --- /dev/null +++ b/core/exceptions.py @@ -0,0 +1,31 @@ +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") + + +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") 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/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/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/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/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/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 b1b5695..f895251 100644 --- a/main.py +++ b/main.py @@ -1,52 +1,274 @@ -from utils.do_logging import logger +import asyncio +import logging +import time import uvicorn -import config -from utils.pre_start_init import app -import routes, models -from fastapi.openapi.utils import get_openapi - - -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": "your@email.com"}, +from config import settings +import os +import gc +from contextlib import asynccontextmanager +from fastapi import FastAPI, Request +from fastapi.middleware.cors import CORSMiddleware +from fastapi.staticfiles import StaticFiles +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 +from VoiceActivityDetector import vad + +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 + +logger = logging.getLogger(__name__) + + +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 + message["headers"] = headers.raw + await send(message) + + try: + await self.app(scope, receive, send_with_request_id) + finally: + 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 {"/", "/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 + await send(message) + + await self.app(scope, receive, send_with_deprecation) + + +@asynccontextmanager +async def lifespan(app): + # Настройка логирования до любых других операций + setup_logging() + + # Установка HF_HOME для HuggingFace Hub + os.environ["HF_HOME"] = settings.HF_HOME + + # on_start + 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") + + # Запуск фоновой push-рассылки статуса (Задача 6.6) + app.state.ws_manager.start_status_broadcast( + metrics_collector=app.state.metrics_collector, + interval_sec=settings.WS_STATUS_BROADCAST_INTERVAL_SEC, ) - # Добавляем пример для 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 + # Запуск фоновой задачи записи метрик в БД (Этап 5) + app.state.metrics_task = asyncio.create_task( + metrics_reporter_loop(app.state, interval_sec=30.0) + ) + + # Настройка сборщика мусора. + 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) + # Инициируем 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("Диаризация недоступна: модель не найдена и не удалось скачать") + + # Инициируем потоковый движок T-one заранее (готов сразу после старта, + # первый коннект на /api/v1/asr/ws-stream не ловит задержку загрузки модели). + try: + import numpy as _np + from Recognizer.tone_engine import get_tone_pipeline + app.state.tone_pipeline = get_tone_pipeline() + # Прогрев инференса одним кадром тишины (аллокация/JIT), чтобы первый чанк был быстрым + _warm = _np.zeros(settings.TONE_CHUNK_SAMPLES, dtype=_np.int32) + _, _st = app.state.tone_pipeline.forward(_warm, None, is_last=True) + app.state.tone_pipeline.finalize(_st) + logger.info("Потоковый движок T-one готов (прогрет)") + except Exception as exc: + app.state.tone_pipeline = None + logger.error("Не удалось инициализировать T-one на старте: %s", exc) + + 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 + ) + observer_thread.start() + logger.info("File watcher started") + + 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() + await app.state.ws_manager.disconnect_all() + + # cleanup (если нужно) + if hasattr(app.state, "recognizer"): + del app.state.recognizer + if hasattr(app.state, "punctuator"): + del app.state.punctuator + if hasattr(app.state, "diarizer"): + del app.state.diarizer + if hasattr(app.state, "tone_pipeline"): + del app.state.tone_pipeline + + +app = FastAPI( + lifespan=lifespan, + version="1.0", + docs_url='/docs', + title='ASR', + 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, + trusted_hosts=settings.TRUSTED_PROXIES, +) + +# TrustedHost middleware +app.add_middleware( + TrustedHostMiddleware, + allowed_hosts=settings.ALLOWED_HOSTS, +) + +# CORS middleware +app.add_middleware( + CORSMiddleware, + allow_origins=settings.CORS_ORIGINS, + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + +# GZip middleware +app.add_middleware( + GZipMiddleware, + minimum_size=500 +) + +# Rate limiting middleware (Этап 5) +app.add_middleware( + RateLimitMiddleware, + max_requests=60, + window_seconds=60.0, +) + +# Exception handlers +register_exception_handlers(app) + +# Static files +app.mount("/static", StaticFiles(directory="static"), name="static") + +# Routers +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__': - app.openapi = custom_openapi - uvicorn.run(app, host=config.HOST, port=config.PORT) + # app.openapi = app.openapi_schema + uvicorn.run(app, host=settings.HOST, port=settings.PORT) except KeyboardInterrupt: logger.info('\nDone') except Exception as e: logger.error(f'\nDone with error {e}') - 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/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..d4d15f4 --- /dev/null +++ b/models/domain/user.py @@ -0,0 +1,90 @@ +from datetime import datetime +from decimal import Decimal +from typing import Optional + +from pydantic import BaseModel, Field + +from config import settings +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=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 + 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" diff --git a/models/fast_api_models.py b/models/fast_api_models.py index 234efa3..df5deac 100644 --- a/models/fast_api_models.py +++ b/models/fast_api_models.py @@ -1,8 +1,94 @@ -from pydantic import BaseModel, HttpUrl, Field -from typing import Union, Annotated +from pydantic import BaseModel, HttpUrl, Field, ConfigDict, RootModel +from typing import Union, Annotated, Optional, Any, List, Dict from fastapi import UploadFile -import config +from config import settings + + +class BaseResponse(BaseModel): + """ + Базовая модель ответа API. + """ + success: bool = True + error_description: Optional[str] = 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: Any = {} + + +class RawData(RootModel[Dict[str, Any]]): + """Структура сырых данных ASR. Словарь каналов, где каждый канал — список результатов.""" + + +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. + """ + 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 + + +class UserBase(BaseModel): + """ + Базовая информация о пользователе (заготовка для JWT). + """ + username: Optional[str] = None + email: Optional[str] = None + is_active: bool = True class SyncASRRequest(BaseModel): @@ -22,12 +108,29 @@ 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={ + "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): @@ -45,82 +148,27 @@ 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 -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 - -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 - } - """ - pass \ No newline at end of file + 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 + } + } + ) 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/models/ws_models.py b/models/ws_models.py new file mode 100644 index 0000000..0bbcbeb --- /dev/null +++ b/models/ws_models.py @@ -0,0 +1,156 @@ +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 + do_dialogue: bool = False + do_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 + 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 + uptime_formatted: str = "0s" + temperature_celsius: float | None = None + + +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/readme.md b/readme.md index 8178fe5..5170cf7 100644 --- a/readme.md +++ b/readme.md @@ -238,6 +238,7 @@ sudo systemctl enable vosk_gpu свяжитесь со мной через [GitHub Issues](https://github.com/Sanich137/ASR_FastAPI_WS_RU_sherpa-onnx/issues) или [GitHub Discussions](https://github.com/Sanich137/ASR_FastAPI_WS_RU_sherpa-onnx/discussions) ## Работы +- 12 мая 2026 года Реализован веб интерфейс. /asr для пользователя и /admin/login для администратора. (без функций реального управления) - 24 декабря 2025 года - обновлён и объединён в один докерфайл. - 11 декабря 2025 года - добавлена батчинг и поддержка RTX для ctc и частично rnnt моделей. - 05 декабря 2025 год - миграция с sherpa-onnx на onnx-asr. diff --git a/requirements.txt b/requirements.txt index c00d457..a6b794a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,6 @@ +pytest setuptools +pydantic_settings httpx~=0.28.1 pathlib~=1.0.1 @@ -20,8 +22,6 @@ tqdm~=4.67.1 psutil~=7.0.0 - - librosa~=0.11.0 # For punctuation @@ -48,3 +48,15 @@ watchdog~=6.0.0 onnx-asr == 0.10.2 hf_xet onnxruntime + +starlette + +# For Auth / JWT +pyjwt~=2.10.1 +bcrypt~=4.3.0 +email-validator + +# For DB (новый api/v1 + сессии ASR; апстрим забыл их в requirements) +sqlalchemy~=2.0 +aiosqlite +greenlet 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/admin.py b/routes/admin.py new file mode 100644 index 0000000..4bf01da --- /dev/null +++ b/routes/admin.py @@ -0,0 +1,128 @@ +"""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("/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, + _admin=Depends(require_admin), +): + """Страница мониторинга сессий.""" + 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, + _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/is_alive.py b/routes/is_alive.py deleted file mode 100644 index bdd2aa4..0000000 --- a/routes/is_alive.py +++ /dev/null @@ -1,45 +0,0 @@ -from utils.pre_start_init import app -import logging -import datetime -import os -import pynvml -from utils.pre_start_init import audio_to_asr - - -def get_gpu_free_memory(): - try: - pynvml.nvmlInit() - handle = pynvml.nvmlDeviceGetHandleByIndex(0) # Первая видеокарта - 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 - - -@app.get("/is_alive") -async def check_if_service_is_alive(): - - logging.info('GET_is_alive') - tasks_in_work = len(audio_to_asr) - - free_mb, gpu_load,temperature = get_gpu_free_memory() - - if tasks_in_work == 0: - state = "idle" - else: - state = "in_work" - - return {"error": False, - "error_description": None, - "state": state, - "tasks_in_work": tasks_in_work, - "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_url.py b/routes/post_by_url.py deleted file mode 100644 index e05e516..0000000 --- a/routes/post_by_url.py +++ /dev/null @@ -1,63 +0,0 @@ -import uuid -import asyncio -import os -from utils.pre_start_init import app, 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 Recognizer.engine.file_recognition import process_file -from threading import Lock -from io import BytesIO - - -# Глобальный лок для потокобезопасности -audio_lock = Lock() - -@app.post("/post_one_step_req") -async def post(params: SyncASRRequest): - """ - На вход принимает HttpUrl - прямую ссылку на скачивание файла 'mp3', 'wav' или 'ogg'.\n - Если на вход передаётся не моно, то ответ будет в несколько элементов списка для каждого канала.\n - По умолчанию отдаёт сырой результат распознавания с разбивкой на части продолжительностью около 15 секунд\n - - :param: do_dialogue: - true, если нужно разбить речь на диалог\n - :param: do_punctuation - true, если нужно расставить пунктуацию. Применяется к диалогу, общему тексту. В проекте.\n - При проектировании таймаутов учитывайте скорость распознавания (около 100 секунд аудио распознаётся за 2-5 секунд - распознавания одного канала) - """ - res = True - error_description = str() - - result = { - "success": res, - "error_description": error_description, - "raw_data": dict(), - "sentenced_data": dict(), - } - - # Получаем файл - 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 { - "success": False, - "error_description": error_description, - "raw_data": dict(), - "sentenced_data": dict(), - } - - try: - # Запускаем обработку в потоке - result = await asyncio.to_thread(process_file, posted_and_downloaded_audio[post_id], params) - except Exception as e: - error_description = f"Ошибка обработки в process_file - {e}" - logger.error(error_description) - result["success"] = False - result['error_description'] = str(error_description) - - return result \ No newline at end of file diff --git a/routes/post_ws.py b/routes/post_ws.py deleted file mode 100644 index 6c674f4..0000000 --- a/routes/post_ws.py +++ /dev/null @@ -1,7 +0,0 @@ -from utils.pre_start_init import app -from models.fast_api_models import WebSocketModel - -@app.post("/ws") -async def post_not_websocket(ws:WebSocketModel): - """Описание для вебсокета ниже в описании WebSocketModel """ - return f"Прочти инструкцию в Schemas - 'WebSocketModel'" diff --git a/routes/root.py b/routes/root.py deleted file mode 100644 index f8749b9..0000000 --- a/routes/root.py +++ /dev/null @@ -1,11 +0,0 @@ -from utils.pre_start_init import app -import config - -@app.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 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/routes/ws_audio_transkrib.py b/routes/ws_audio_transkrib.py index a353802..2b1bc11 100644 --- a/routes/ws_audio_transkrib.py +++ b/routes/ws_audio_transkrib.py @@ -1,39 +1,55 @@ from pydub import AudioSegment import ujson -import config +import base64 +import logging +from config import settings import uuid from io import BytesIO -import subprocess - -from utils.pre_start_init import app -from fastapi import WebSocket, WebSocketException -from utils.do_logging import logger +from fastapi import APIRouter, WebSocket +from services.ws_protocol import normalize_to_ws_message +from models.ws_models import WSConfigMessage, WSAudioMessage, WSEosMessage 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 + +# Todo: Этот роут — legacy. Удалить после полного перехода клиентов на /api/v1/asr/ws. +# Сохраняем "frozen" для обратной совместимости; не использовать в новых интеграциях. + +router = APIRouter() +logger = logging.getLogger(__name__) -@app.websocket("/ws") -async def websocket(ws: WebSocket): +@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=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 @@ -49,35 +65,40 @@ async def websocket(ws: WebSocket): 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')}") + # Автодетект протокола: понимает legacy ({config}/{eof}) и новый ({type:...}) + canonical = normalize_to_ws_message(message.get('text')) + if isinstance(canonical, WSConfigMessage): + audio_format = canonical.audio_format + sample_rate = canonical.sample_rate + wait_null_answers = canonical.wait_null_answers + do_dialogue = canonical.do_dialogue + do_punctuation = canonical.do_punctuation + channel_name = canonical.channel_name or "Null" + logger.info(f"Task received, config - sr={sample_rate}, fmt={audio_format}, ch={channel_name}") continue - - elif message.get('text') and 'eof' in message.get('text'): - logger.info(f"EOF received in channel {channel_name}") + elif isinstance(canonical, WSEosMessage): + logger.info(f"EOF/EOS received in channel {channel_name}") break + elif isinstance(canonical, WSAudioMessage) and canonical.audio_base64: + # Новый протокол с base64-аудио -> превращаем в бинарный кадр, обрабатываем ниже + message = {'bytes': base64.b64decode(canonical.audio_base64)} else: - logger.error(f"Can`t recognise text part of message {message.get('text')} in channel {channel_name}") - + logger.error(f"Can`t recognise text message {message.get('text')} in channel {channel_name}") + continue 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'): + continue + + if 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, @@ -85,6 +106,7 @@ async def websocket(ws: WebSocket): sample_width = 2, # Ширина сэмпла (2 байта для int16) channels = 1 # Количество каналов. По умолчанию - 1, Моно. ) + else: try: buffer = BytesIO(chunk) @@ -98,26 +120,26 @@ 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) - # Копим буфер 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: 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]) + 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) @@ -191,7 +213,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}') @@ -215,7 +237,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/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/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 new file mode 100644 index 0000000..c68613c --- /dev/null +++ b/services/asr_pipeline.py @@ -0,0 +1,326 @@ +""" +Модуль services/asr_pipeline.py +Содержит бизнес-логику потокового распознавания речи (ASR) через WebSocket: +накопление аудио, VAD-разделение по паузам, распознавание, постпроцессинг, +накопление результатов и формирование финального диалога с пунктуацией. +""" + +import logging +from io import BytesIO +from typing import Union, Annotated, Optional, Any, List, Dict + +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, + metrics_collector: Optional[Any] = None, +) -> 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 + 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 + + # --- 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 + + # 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 + + segment = session.audio_to_asr[-1] + if segment.duration_seconds <= 0: + return + + 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) + finally: + 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 == " " + + 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, + metrics_collector: Optional[Any] = None, +) -> 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-соединений. + """ + if metrics_collector is not None: + metrics_collector.increment_tasks() + try: + # --- 1. Объединение остатков --- + final_audio = session.audio_overlap + session.audio_buffer + 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: + 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. Распознавание финального остатка --- + 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() + + # --- 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 + 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/recognition_session.py b/services/recognition_session.py new file mode 100644 index 0000000..063fea7 --- /dev/null +++ b/services/recognition_session.py @@ -0,0 +1,163 @@ +""" +Модуль 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" + connecting = "connecting" + receiving = "receiving" + 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 + + async 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_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..b8d0ad2 --- /dev/null +++ b/services/ws_manager.py @@ -0,0 +1,245 @@ +""" +Модуль 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 + user_id: Optional[str] = None + + +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] = {} + 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: + """Текущее количество активных соединений.""" + 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) + + 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: + """ + Принудительно закрывает все активные соединения. + + Используется при 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..e0e9f4c --- /dev/null +++ b/services/ws_metrics.py @@ -0,0 +1,216 @@ +""" +Модуль 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], Optional[float]]: + """ + Возвращает свободную и общую память RAM в мегабайтах и загрузку CPU в процентах. + + Returns: + Кортеж (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) + 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: + """Увеличивает счётчик активных задач распознавания на 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. + - GPU utilization > 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() + 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 + 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: + 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" + + @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, + max_connections: int = 100, + ) -> WSStatusResponse: + """ + Собирает полный набор метрик и возвращает WSStatusResponse. + + Args: + active_connections: Текущее количество активных WS-соединений. + max_connections: Максимально допустимое количество соединений. + + Returns: + WSStatusResponse с актуальными метриками. + """ + gpu_free, gpu_total = self.get_gpu_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) + + return WSStatusResponse( + 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), + uptime_formatted=self._format_uptime(uptime), + temperature_celsius=gpu_temp, + ) diff --git a/services/ws_protocol.py b/services/ws_protocol.py new file mode 100644 index 0000000..1f0a422 --- /dev/null +++ b/services/ws_protocol.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +""" +Автодетект и нормализация входящих WebSocket-сообщений. + +Поддерживает ДВА протокола в одном эндпоинте: + - legacy (asterisk-socket-server): {"config": {...}, "channelName": ...} + сырые PCM16 байты (binary frame) + {"eof": 1} + - новый (api/v1): {"type": "config", "sample_rate": ..., ...} + {"type": "audio_chunk", "audio_base64": ...} | binary frame + {"type": "eos"} / {"type": "ping"} + +Возвращает единое событие WSEvent, чтобы эндпоинт не зависел от формата. +""" + +import base64 +from dataclasses import dataclass +from typing import Optional + +import ujson + +from models.ws_models import ( + parse_ws_message, + WSBaseMessage, + WSConfigMessage, + WSEosMessage, +) + + +@dataclass +class WSEvent: + kind: str # "config" | "audio" | "eos" | "ping" | "disconnect" | "ignore" + audio: Optional[bytes] = None + sample_rate: Optional[int] = None + channel_name: Optional[str] = None + wait_null_answers: bool = False + raw: Optional[dict] = None # исходный распарсенный JSON (для расширений) + + +def detect(message: dict) -> WSEvent: + """ + Нормализует одно сообщение из ws.receive() (dict с ключами type/text/bytes). + """ + if message.get("type") == "websocket.disconnect": + return WSEvent("disconnect") + + # Бинарный фрейм аудио — одинаков в обоих протоколах + raw_bytes = message.get("bytes") + if raw_bytes: + return WSEvent("audio", audio=raw_bytes) + + text = message.get("text") + if not text: + return WSEvent("ignore") + + try: + d = ujson.loads(text) + except Exception: + return WSEvent("ignore", raw=None) + + if not isinstance(d, dict): + return WSEvent("ignore") + + # --- legacy --- + if isinstance(d.get("config"), dict): + cfg = d["config"] + return WSEvent( + "config", + sample_rate=cfg.get("sample_rate"), + channel_name=d.get("channelName") or cfg.get("channelName"), + wait_null_answers=bool(cfg.get("wait_null_answers", False)), + raw=d, + ) + if "eof" in d or "eos" in d: + return WSEvent("eos", raw=d) + + # --- новый протокол (дискриминатор type) --- + mtype = d.get("type") + if mtype == "config": + return WSEvent( + "config", + sample_rate=d.get("sample_rate"), + channel_name=d.get("channel_name"), + wait_null_answers=bool(d.get("wait_null_answers", False)), + raw=d, + ) + if mtype == "audio_chunk": + b64 = d.get("audio_base64") + try: + audio = base64.b64decode(b64) if b64 else b"" + except Exception: + audio = b"" + return WSEvent("audio", audio=audio, raw=d) + if mtype in ("eos", "eof"): + return WSEvent("eos", raw=d) + if mtype == "ping": + return WSEvent("ping", raw=d) + + return WSEvent("ignore", raw=d) + + +def normalize_to_ws_message(text: str | bytes) -> Optional[WSBaseMessage]: + """ + Приводит текстовое WS-сообщение ЛЮБОГО протокола к каноническому pydantic-объекту + (WSConfigMessage / WSEosMessage / WSAudioMessage / ...). Понимает: + - новый протокол (дискриминатор type) — через parse_ws_message; + - legacy ({"config": {...}, "channelName": ...} и {"eof": 1}). + Возвращает None, если распознать не удалось. + + Используется в офлайн-эндпоинтах (/api/v1/asr/ws и legacy /ws), чтобы они принимали + оба протокола без изменения нижележащей логики. + """ + # 1) новый протокол + try: + return parse_ws_message(text) + except Exception: + pass + + # 2) legacy + try: + d = ujson.loads(text) + except Exception: + return None + if not isinstance(d, dict): + return None + + if isinstance(d.get("config"), dict): + cfg = d["config"] + try: + return WSConfigMessage( + sample_rate=int(cfg.get("sample_rate") or 16000), + wait_null_answers=bool(cfg.get("wait_null_answers", True)), + do_dialogue=bool(cfg.get("do_dialogue", False)), + do_punctuation=bool(cfg.get("do_punctuation", False)), + audio_format=cfg.get("audio_format", "pcm16"), + channel_name=d.get("channelName") or cfg.get("channelName"), + ) + except Exception: + return None + if "eof" in d or "eos" in d: + return WSEosMessage() + return None diff --git a/services/ws_session.py b/services/ws_session.py new file mode 100644 index 0000000..4d1cb54 --- /dev/null +++ b/services/ws_session.py @@ -0,0 +1,191 @@ +""" +Модуль services/ws_session.py +Содержит класс AudioSession для управления состоянием и буфером аудио +в рамках одной WebSocket-сессии распознавания речи (ASR). +""" + +import time +from collections import deque +from typing import Optional + +import numpy as np + +from models.ws_models import WSConfigMessage +from services.recognition_session import RecognitionSession, SessionState +from config import settings + + +class AudioSession(RecognitionSession): + """ + Управляет буфером аудио и состоянием для одного 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. + """ + 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 = ( + max_buffer_duration_sec + if max_buffer_duration_sec is not None + else getattr(settings, "WS_MAX_BUFFER_DURATION_SEC", 300.0) + ) + self.buffer: deque[np.ndarray] = deque() + + @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: + """ + Вычисляет суммарную длительность аудио в буфере (в секундах). + + Использует 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: + """ + Очищает WS-специфичные буферы и сбрасывает базовое состояние. + """ + await super().reset() + self.state = SessionState.connecting + self.buffer.clear() + self.last_activity = time.time() + self.wait_null_answers = True + + 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 + + def to_dict(self) -> dict: + """ + Сериализует лёгкие (не AudioSegment) поля сессии в dict для StateStore. + + Returns: + dict с client_id, config, channel_name, audio_duration, флагами + и накопленными результатами распознавания. + """ + data = super().to_dict() + data.update({ + "client_id": self.client_id, + "wait_null_answers": self.wait_null_answers, + "last_activity": self.last_activity, + "max_buffer_duration_sec": self.max_buffer_duration_sec, + }) + return data + + @classmethod + def from_dict(cls, data: dict) -> "AudioSession": + """ + Восстанавливает сессию из dict (только мета-поля, без AudioSegment-буферов). + + Args: + data: dict, полученный из to_dict(). + + Returns: + AudioSession с восстановленной конфигурацией и флагами. + """ + 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") + 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.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/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/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..366b9d9 --- /dev/null +++ b/static/js/admin_ui.js @@ -0,0 +1,55 @@ +(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; + } + + function escapeHtml(str) { + if (!str) return ''; + const div = document.createElement('div'); + div.textContent = str; + return div.innerHTML; + } + + window.AdminUI = { renderBadge, renderDate, renderDuration, escapeHtml, initAdmin }; +})(); 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..a915015 --- /dev/null +++ b/static/js/asr_page.js @@ -0,0 +1,755 @@ +(function() { + 'use strict'; + + let wsChannelResults = []; + let micRecorder = null; + let micChunks = []; + let micBlob = null; + let micStream = null; + let micStartTime = null; + let micTimerInterval = null; + + function extractResultData(payload) { + const payloads = Array.isArray(payload) ? payload : [payload]; + let hasDialog = false; + let dialogText = ''; + let textContent = ''; + let rawJson = ''; + + payloads.forEach((p, idx) => { + const prefix = payloads.length > 1 ? `=== Канал ${idx+1} ===\n` : ''; + const sentenced = p.sentenced_data || p.data?.sentenced_data; + + if (idx > 0) { + if (dialogText) dialogText += '\n\n'; + if (textContent) textContent += '\n\n'; + if (rawJson) rawJson += '\n\n'; + } + + rawJson += prefix + JSON.stringify(p, null, 2); + + if (sentenced?.full_text_only) { + const ft = sentenced.full_text_only; + textContent += prefix + (Array.isArray(ft) ? ft.join('\n') : String(ft)); + } else { + textContent += prefix + (sentenced?.raw_text_sentenced_recognition || p.data?.text || p.data?.raw_data?.channel_1?.map(x => x.data?.text).join('\n') || ''); + } + + const items = sentenced?.list_of_sentenced_recognitions; + if (items && items.length > 0) { + hasDialog = true; + if (dialogText) dialogText += '\n\n'; + dialogText += prefix + items.map(item => { + const start = item.start != null ? item.start : (item.start_time || ''); + const speaker = item.speaker != null ? `[${item.speaker}]` : ''; + return `${speaker} ${start} - ${item.text || ''}`.trim(); + }).join('\n'); + } + + const diarized = p.diarized_data || p.data?.diarized_data; + if (diarized && Array.isArray(diarized) && diarized.length > 0) { + hasDialog = true; + if (dialogText) dialogText += '\n\n'; + dialogText += prefix + diarized.map(item => { + const start = item.start != null ? item.start : (item.start_time || ''); + const speaker = item.speaker != null ? `[${item.speaker}]` : ''; + return `${speaker} ${start} - ${item.text || ''}`.trim(); + }).join('\n'); + } + }); + + return { hasDialog, dialogText, textContent, rawJson }; + } + + function displayResult(mode, payload) { + const data = extractResultData(payload); + const container = document.getElementById(mode + 'Result'); + const pre = document.getElementById(mode + 'ResultText'); + if (!container || !pre) return; + + container.style.display = 'block'; + + let btnContainer = document.getElementById(mode + 'FormatButtons'); + if (!btnContainer) { + const card = pre.parentElement; + btnContainer = document.createElement('div'); + btnContainer.className = 'flex gap-2 mb-2'; + btnContainer.id = mode + 'FormatButtons'; + card.insertBefore(btnContainer, pre); + } + btnContainer.innerHTML = ` + + + + `; + + window._resultData = window._resultData || {}; + window._resultData[mode] = data; + + const defaultFormat = data.hasDialog ? 'dialog' : 'text'; + switchResultFormat(mode, defaultFormat); + } + + function highlightJson(json) { + if (typeof json !== 'string') json = JSON.stringify(json, null, 2); + return json.replace(/&/g, '&').replace(//g, '>') + .replace(/("(\\u[a-zA-Z0-9]{4}|\\[^u]|[^\\"])*"(\s*:)?)/g, function(match) { + let cls = 'json-string'; + if (/:$/.test(match)) { cls = 'json-key'; match = match.slice(0, -1) + ':'; return '' + match; } + return '' + match + ''; + }) + .replace(/\b(true|false|null)\b/g, '$1') + .replace(/\b(\d+\.?\d*)\b/g, '$1'); + } + + window.switchResultFormat = function(mode, format) { + const data = window._resultData?.[mode]; + if (!data) return; + + const pre = document.getElementById(mode + 'ResultText'); + const btnDialog = document.getElementById(mode + '_btn_dialog'); + const btnText = document.getElementById(mode + '_btn_text'); + const btnRaw = document.getElementById(mode + '_btn_raw'); + + [btnDialog, btnText, btnRaw].forEach(btn => { + if (btn) { + btn.style.opacity = '0.5'; + btn.style.borderColor = 'transparent'; + } + }); + + const activeBtn = { dialog: btnDialog, text: btnText, raw: btnRaw }[format]; + if (activeBtn) { + activeBtn.style.opacity = '1'; + activeBtn.style.borderColor = 'var(--color-primary)'; + } + + if (btnDialog && !data.hasDialog) { + btnDialog.disabled = true; + btnDialog.style.opacity = '0.3'; + btnDialog.style.cursor = 'not-allowed'; + } else if (btnDialog) { + btnDialog.disabled = false; + btnDialog.style.cursor = 'pointer'; + } + + data.currentFormat = format; + if (format === 'dialog') { + pre.textContent = data.dialogText; + } else if (format === 'text') { + pre.textContent = data.textContent; + } else { + pre.innerHTML = highlightJson(data.rawJson); + } + }; + + window.downloadResult = function(mode, btn) { + UI.setLoading(btn, true); + const data = window._resultData?.[mode]; + if (!data) { UI.setLoading(btn, false); return; } + const format = data.currentFormat || 'text'; + let content, filename, mime; + if (format === 'raw') { + content = data.rawJson; + filename = 'result.json'; + mime = 'application/json'; + } else { + content = format === 'dialog' ? data.dialogText : data.textContent; + filename = 'result.txt'; + mime = 'text/plain'; + } + const blob = new Blob([content], {type: mime}); + const a = document.createElement('a'); + a.href = URL.createObjectURL(blob); + a.download = filename; + a.click(); + setTimeout(() => UI.setLoading(btn, false), 500); + }; + + function refreshWsResult() { + if (wsChannelResults.length === 0) return; + const payloads = wsChannelResults.map(c => c.payload); + displayResult('ws', payloads); + } + + // --- Микрофон --- + function startRecording() { + if (!navigator.mediaDevices || !navigator.mediaDevices.getUserMedia) { + UI.toast('Ваш браузер не поддерживает запись с микрофона', 'error'); + return; + } + navigator.mediaDevices.enumerateDevices().then(devices => { + const audioInputs = devices.filter(d => d.kind === 'audioinput'); + if (audioInputs.length === 0) { + UI.toast('Микрофон не обнаружен. Проверьте подключение аудио-устройства и разрешения браузера.', 'error'); + return; + } + navigator.mediaDevices.getUserMedia({ audio: { echoCancellation: false, noiseSuppression: false, autoGainControl: false } }).then(stream => { + micStream = stream; + micChunks = []; + micRecorder = new MediaRecorder(stream); + micRecorder.ondataavailable = e => { if (e.data.size > 0) micChunks.push(e.data); }; + micRecorder.onstop = () => { + micBlob = new Blob(micChunks, { type: 'audio/webm' }); + const url = URL.createObjectURL(micBlob); + const player = document.getElementById('micAudioPlayer'); + player.src = url; + document.getElementById('micRecordingBlock').style.display = 'none'; + document.getElementById('micPreviewBlock').style.display = 'block'; + document.getElementById('btnMicSend').disabled = false; + clearInterval(micTimerInterval); + micTimerInterval = null; + }; + micRecorder.start(); + micStartTime = Date.now(); + document.getElementById('micIdleBlock').style.display = 'none'; + document.getElementById('micRecordingBlock').style.display = 'block'; + document.getElementById('micRecDot').classList.add('mic-recording'); + micTimerInterval = setInterval(() => { + const sec = Math.floor((Date.now() - micStartTime) / 1000); + const mm = String(Math.floor(sec / 60)).padStart(2, '0'); + const ss = String(sec % 60).padStart(2, '0'); + document.getElementById('micTimer').textContent = mm + ':' + ss; + }, 1000); + }).catch(err => { + let msg = 'Не удалось получить доступ к микрофону'; + if (err.name === 'NotFoundError' || err.name === 'DevicesNotFoundError') { + msg = 'Микрофон не найден. Убедитесь, что устройство подключено и браузеру разрешён доступ.'; + } else if (err.name === 'NotAllowedError' || err.name === 'PermissionDeniedError') { + msg = 'Доступ к микрофону запрещён. Разрешите использование микрофона в настройках браузера.'; + } else if (err.name === 'NotReadableError' || err.name === 'TrackStartError') { + msg = 'Микрофон занят другим приложением. Закройте другие программы, использующие микрофон.'; + } else if (err.name === 'SecurityError') { + msg = 'Доступ к микрофону заблокирован. Используйте HTTPS или localhost.'; + } + UI.toast(msg + ' (' + err.message + ')', 'error'); + }); + }).catch(err => { + UI.toast('Не удалось проверить аудио-устройства: ' + err.message, 'error'); + }); + } + window.startRecording = startRecording; + + function stopRecording() { + if (micRecorder && micRecorder.state !== 'inactive') { + micRecorder.stop(); + } + if (micStream) { + micStream.getTracks().forEach(t => t.stop()); + } + document.getElementById('micRecDot').classList.remove('mic-recording'); + } + window.stopRecording = stopRecording; + + function resetRecording() { + if (micRecorder && micRecorder.state !== 'inactive') { + micRecorder.stop(); + } + if (micStream) { + micStream.getTracks().forEach(t => t.stop()); + } + micRecorder = null; + micStream = null; + micBlob = null; + micChunks = []; + if (micTimerInterval) clearInterval(micTimerInterval); + micTimerInterval = null; + document.getElementById('micAudioPlayer').src = ''; + document.getElementById('micIdleBlock').style.display = 'block'; + document.getElementById('micRecordingBlock').style.display = 'none'; + document.getElementById('micPreviewBlock').style.display = 'none'; + document.getElementById('btnMicSend').disabled = true; + document.getElementById('micTimer').textContent = '00:00'; + document.getElementById('micRecDot').classList.remove('mic-recording'); + } + window.resetRecording = resetRecording; + + function getMicParams() { + const expert = document.getElementById('mic_expert').checked; + const defaults = ASRSettings.getDefaults('mic'); + 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('mic_fast_speech'); + const splitPhrases = domBool('mic_split_phrases'); + const form = new FormData(); + form.append('keep_raw', expert ? domBool('mic_keep_raw') : defaults.keep_raw); + form.append('do_echo_clearing', expert ? domBool('mic_do_echo_clearing') : defaults.do_echo_clearing); + form.append('do_dialogue', splitPhrases ? true : (expert ? domBool('mic_do_dialogue') : defaults.do_dialogue)); + form.append('do_diarization', expert ? domBool('mic_do_diarization') : defaults.do_diarization); + form.append('do_punctuation', splitPhrases ? true : (expert ? domBool('mic_do_punctuation') : defaults.do_punctuation)); + form.append('make_mono', expert ? domBool('mic_make_mono') : defaults.make_mono); + form.append('diar_vad_sensity', expert ? domInt('mic_diar_vad_sensity', defaults.diar_vad_sensity) : defaults.diar_vad_sensity); + form.append('do_auto_speech_speed_correction', fastSpeech ? true : (expert ? domBool('mic_do_auto_speech_speed_correction') : defaults.do_auto_speech_speed_correction)); + form.append('speech_speed_correction_multiplier', expert ? domFloat('mic_speech_speed_correction_multiplier', defaults.speech_speed_correction_multiplier) : defaults.speech_speed_correction_multiplier); + form.append('use_batch', fastSpeech ? false : (expert ? domBool('mic_use_batch') : defaults.use_batch)); + form.append('batch_size', expert ? domInt('mic_batch_size', defaults.batch_size) : defaults.batch_size); + return form; + } + + async function sendMic() { + if (!micBlob) { UI.toast('Запишите аудио перед отправкой', 'warning'); return; } + const btn = document.getElementById('btnMicSend'); + UI.setLoading(btn, true); + if (document.getElementById('mic_expert').checked && !validateExpert('mic')) { + UI.setLoading(btn, false); return; + } + const form = getMicParams(); + form.append('file', micBlob, 'recording.webm'); + try { + const resp = await Auth.apiFetch('/api/v1/asr/file', {method:'POST', body:form}); + const data = await resp.json(); + displayResult('mic', data); + } catch (e) { + UI.toast('Ошибка: ' + e.message, 'error'); + } finally { + UI.setLoading(btn, false); + } + } + window.sendMic = sendMic; + + // --- Табы --- + 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, btn) { + UI.setLoading(btn, true); + const text = document.getElementById(elId).textContent; + navigator.clipboard.writeText(text).then(() => { + UI.toast('Скопировано', 'success'); + UI.setLoading(btn, false); + }).catch(() => UI.setLoading(btn, false)); + } + window.copyText = copyText; + + function downloadText(elId, filename, btn) { + UI.setLoading(btn, true); + 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(); + setTimeout(() => UI.setLoading(btn, false), 500); + } + 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 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(payload) + }); + const data = await resp.json(); + displayResult('url', data); + } 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) { + if (file.size > MAX_FILE_SIZE_MB * 1024 * 1024) { + UI.toast('Файл слишком большой. Максимум ' + MAX_FILE_SIZE_MB + ' МБ', 'error'); + selectedFile = null; + document.getElementById('btnFile').disabled = true; + return; + } + selectedFile = file; + document.getElementById('btnFile').disabled = false; + UI.toast('Файл выбран: ' + file.name, 'info'); + } + 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'; + 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); + try { + const resp = await Auth.apiFetch('/api/v1/asr/file', {method:'POST', body:form}); + const data = await resp.json(); + displayResult('file', data); + } 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]; + if (f && f.size > MAX_FILE_SIZE_MB * 1024 * 1024) { + UI.toast('Файл слишком большой. Максимум ' + MAX_FILE_SIZE_MB + ' МБ', 'error'); + document.getElementById('wsFileInput').value = ''; + document.getElementById('btnWs').disabled = true; + return; + } + document.getElementById('btnWs').disabled = !f; + } + window.onWsFileSelected = onWsFileSelected; + + const MAX_FILE_SIZE_MB = 100; + + function setWsStatus(status) { + const dot = document.getElementById('wsStatusDot'); + const txt = document.getElementById('wsStatusText'); + dot.classList.remove('ws-connecting'); + if (status === 'connected') { dot.style.background = 'var(--color-success)'; txt.textContent = 'Подключено'; } + else if (status === 'connecting') { dot.style.background = 'var(--color-warning)'; txt.textContent = 'Подключение...'; dot.classList.add('ws-connecting'); } + 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); + } + + 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 token = Auth.getAccessToken(); + + socket.onopen = async function() { + 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: 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); + 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) { + wsChannelResults.push({ channel: channel + 1, payload: data }); + refreshWsResult(); + 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, wsParams) { + for (let channel = 0; channel < numChannels; channel++) { + await sendChannel(arrayBuffer, channel, numChannels, chunkSize, sampleRate, wsParams); + console.log("конец канала"); + } + } + + 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 = ''; + document.getElementById('wsPartial').textContent = ''; + document.getElementById('btnWsStop').style.display = 'inline-flex'; + setWsStatus('connecting'); + wsSockets = []; + wsChannelResults = []; + const wsParams = getWsParams(); + + try { + const arrayBuffer = await file.arrayBuffer(); + const { sampleRate, numChannels } = readWavHeader(arrayBuffer); + const chunkSize = 65536; + await sendAllChannels(arrayBuffer, numChannels, chunkSize, sampleRate, wsParams); + 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 => { + const name = s.file_name || (s.audio_url ? s.audio_url.split('/').pop() : null) || '-'; + return ` + + + + + `; + }).join('') + '
${UI.formatDate(s.created_at)}${s.status}${s.session_type}${name}
'; + } catch (e) { + el.innerHTML = '
Не удалось загрузить
'; + } + } + + // 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..73fce0f --- /dev/null +++ b/static/js/asr_settings.js @@ -0,0 +1,115 @@ +(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, + }, + mic: { + 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/static/js/auth.js b/static/js/auth.js new file mode 100644 index 0000000..ae0842f --- /dev/null +++ b/static/js/auth.js @@ -0,0 +1,104 @@ +(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; _setCookie('access_token', t, 1); } + function getAccessToken() { return _token; } + function clearAuth() { _token = null; _setCookie('access_token', '', -1); } + + 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, refreshToken }; +})(); 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..d903090 --- /dev/null +++ b/templates/admin/api_keys.html @@ -0,0 +1,153 @@ +{% extends "admin/base_admin.html" %} +{% block title %}API-ключи{% endblock %} +{% block content %} +

API-ключи

+ +
+
+
+ + +
+ +
+
+ +
+
+
+
+ +
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/base_admin.html b/templates/admin/base_admin.html new file mode 100644 index 0000000..7a25a14 --- /dev/null +++ b/templates/admin/base_admin.html @@ -0,0 +1,80 @@ + + + + + + + {% block title %}Админ-панель{% endblock %} + + + + {% block head %}{% endblock %} + + +
+ ASR Admin + +
+
+ +
+ {% block content %}{% endblock %} +
+
+ + + + + + {% block scripts %}{% endblock %} + + diff --git a/templates/admin/components/tariff_modal.html b/templates/admin/components/tariff_modal.html new file mode 100644 index 0000000..85b09de --- /dev/null +++ b/templates/admin/components/tariff_modal.html @@ -0,0 +1,137 @@ + +
+

Создать тариф

+ + +
+ + + +
+ +
+ + + +
+ +
+ + + +
+ +
+ + +
+ +
+ + +
+ +
+ +
+ +
+ + +
+
+
+ + diff --git a/templates/admin/dashboard.html b/templates/admin/dashboard.html new file mode 100644 index 0000000..1b5996b --- /dev/null +++ b/templates/admin/dashboard.html @@ -0,0 +1,105 @@ +{% extends "admin/base_admin.html" %} + +{% block title %}Dashboard — Админ-панель{% endblock %} + +{% block content %} +
+

Dashboard

+ + + + + +
+
+
CPU
+
+
+
+
GPU
+
+
+
+
Активные задачи
+
+
+
+
Очередь
+
+
+
+
Uptime
+
+
+
+
Adapter Status
+
+
+
+ + +
+
+

Нагрузка за 24 часа

+ Загрузка... +
+ +
+
+ +{% endblock %} + +{% block scripts %} + + +{% endblock %} diff --git a/templates/admin/login.html b/templates/admin/login.html new file mode 100644 index 0000000..d6da272 --- /dev/null +++ b/templates/admin/login.html @@ -0,0 +1,42 @@ + + + + + + Вход в админ-панель + + + +
+

Admin Login

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

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

+ +
+
+
+
+ + +
+
+ + +
+ +
+ +
+
+ +
+
+
+
+ +
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/session_detail.html b/templates/admin/session_detail.html new file mode 100644 index 0000000..24869aa --- /dev/null +++ b/templates/admin/session_detail.html @@ -0,0 +1,302 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Сессия — Админ-панель{% endblock %} +{% block content %} +
+
+ ← Назад +

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

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

Результат

+
+ + +
+
+
+
+
+ + +
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/sessions.html b/templates/admin/sessions.html new file mode 100644 index 0000000..6b665f0 --- /dev/null +++ b/templates/admin/sessions.html @@ -0,0 +1,102 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Сессии{% endblock %} +{% block content %} +

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

+ +
+
+
Последнее обновление: —
+ +
+
+
+
+
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/settings.html b/templates/admin/settings.html new file mode 100644 index 0000000..a6a3bfd --- /dev/null +++ b/templates/admin/settings.html @@ -0,0 +1,61 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Настройки{% endblock %} +{% block content %} +

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

+ +
+

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

+

При включении режима обслуживания новые ASR-запросы будут отклонены с кодом 503.

+
+ + Загрузка... +
+
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/subscriptions.html b/templates/admin/subscriptions.html new file mode 100644 index 0000000..f5538e9 --- /dev/null +++ b/templates/admin/subscriptions.html @@ -0,0 +1,202 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Подписки{% endblock %} +{% block content %} +

Подписки

+ +
+
+
+ + +
+
+ + +
+ +
+
+ +
+
+
+
+ +
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/tariffs.html b/templates/admin/tariffs.html new file mode 100644 index 0000000..738945d --- /dev/null +++ b/templates/admin/tariffs.html @@ -0,0 +1,104 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Тарифы{% endblock %} +{% block content %} +
+
+

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

+ +
+ +
+ + + + + + + + + + + + + + + +
КодНазваниеЗапросов/минАудио (сек)Цена/месСтатусДействия
+
+
+
+
+ +
+
+ +{% include "admin/components/tariff_modal.html" %} +{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/telegram.html b/templates/admin/telegram.html new file mode 100644 index 0000000..8f93c24 --- /dev/null +++ b/templates/admin/telegram.html @@ -0,0 +1,98 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Telegram{% endblock %} +{% block content %} +

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

+ +
+

Настройки

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

Статистика

+
+
+
+
+{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/admin/transactions.html b/templates/admin/transactions.html new file mode 100644 index 0000000..a8e3900 --- /dev/null +++ b/templates/admin/transactions.html @@ -0,0 +1,156 @@ +{% 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 new file mode 100644 index 0000000..22fbfb5 --- /dev/null +++ b/templates/admin/users.html @@ -0,0 +1,132 @@ +{% extends "admin/base_admin.html" %} +{% block title %}Пользователи{% endblock %} +{% block content %} +
+
+

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

+
+ +
+
+ +
+ + + + +
+ + + +
+
+ +
+ + + + + + + + + + + + + + + +
EmailИмяРольСтатусTelegramРегистрацияПоследний вход
+
+
+
+
+ +
+
+ + +{% endblock %} + +{% block scripts %} + +{% endblock %} diff --git a/templates/auth/login.html b/templates/auth/login.html new file mode 100644 index 0000000..ad1e512 --- /dev/null +++ b/templates/auth/login.html @@ -0,0 +1,69 @@ +{% 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..e20c3b0 --- /dev/null +++ b/templates/auth/register.html @@ -0,0 +1,75 @@ +{% 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/index.html b/templates/index.html index 006ef5c..9e472aa 100644 --- a/templates/index.html +++ b/templates/index.html @@ -65,6 +65,10 @@

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

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

serverResponses.innerHTML = '

Загрузка...

'; - fetch('/post_one_step_req', { + fetch('/api/v1/asr/url', { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -252,28 +256,43 @@

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

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}/ws`; + const wsUrl = `${wsProtocol}//${window.location.host}/api/v1/asr/ws`; const socket = new WebSocket(wsUrl); const doDialogue = document.getElementById('ws_do_dialogue').checked; 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) { @@ -367,7 +395,7 @@

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

try { serverResponses.innerHTML = '

Загрузка...

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

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

- \ No newline at end of file + 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..aca40a9 --- /dev/null +++ b/templates/user/asr.html @@ -0,0 +1,613 @@ +{% 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..cd911ab --- /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..cb00bd7 --- /dev/null +++ b/templates/user/dashboard.html @@ -0,0 +1,140 @@ +{% 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..21e791e --- /dev/null +++ b/templates/user/history.html @@ -0,0 +1,211 @@ +{% 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..e3c3133 --- /dev/null +++ b/templates/user/history_detail.html @@ -0,0 +1,273 @@ +{% 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..15ab7b6 --- /dev/null +++ b/templates/user/profile.html @@ -0,0 +1,153 @@ +{% 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..446031f --- /dev/null +++ b/templates/user/subscription.html @@ -0,0 +1,140 @@ +{% extends "user/base_user.html" %} + +{% block title %}Подписка — ASR Сервис{% endblock %} + +{% block user_content %} +
+

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

+ + +
+
+
+ + +

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

+
+
+
+ + +

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

+
+
+
+
+ + +{% endblock %} diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 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_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..0f7641d --- /dev/null +++ b/tests/test_gzip_middleware.py @@ -0,0 +1,67 @@ +import pytest +from fastapi.testclient import TestClient +from starlette.middleware.gzip import GZipMiddleware +from main import app + + +@pytest.fixture +def client(): + # Сбрасываем кэш middleware stack, чтобы перестроить с актуальными параметрами + app.middleware_stack = None + 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"}) + 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 + 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_http_routes.py b/tests/test_http_routes.py new file mode 100644 index 0000000..c73b178 --- /dev/null +++ b/tests/test_http_routes.py @@ -0,0 +1,123 @@ +import asyncio +import json + +import httpx +import pytest + +from config import settings +from models.fast_api_models import V1BaseResponse as BaseResponse + +BASE_URL = f"http://127.0.0.1:{settings.PORT}/api/v1" + + +async def _assert_base_response(body: dict, expect_success: bool): + assert "success" in body, f"Ответ должен содержать поле 'success' (V1BaseResponse)" + assert body["success"] is expect_success + + +@pytest.mark.asyncio +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, + } + try: + resp = await client.post( + f"{BASE_URL}/asr/url", 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)) + + await _assert_base_response(body, expect_success=True) + 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 + try: + resp = await client.post( + f"{BASE_URL}/asr/url", 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)) + + assert resp.status_code == 422 + await _assert_base_response(body, expect_success=False) + assert body.get("error_description") is not None + 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: + 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", + } + try: + resp = await client.post( + f"{BASE_URL}/asr/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)) + + await _assert_base_response(body, expect_success=True) + 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"} + try: + resp = await client.post( + f"{BASE_URL}/asr/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)) + + 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()) 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/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_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_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_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" 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_trusted_host_middleware.py b/tests/test_trusted_host_middleware.py new file mode 100644 index 0000000..fdf61ac --- /dev/null +++ b/tests/test_trusted_host_middleware.py @@ -0,0 +1,35 @@ +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.""" + 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 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..4e316f4 --- /dev/null +++ b/tests/test_ws_models.py @@ -0,0 +1,259 @@ +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_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) + 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..c76c3f6 --- /dev/null +++ b/tests/test_ws_session.py @@ -0,0 +1,141 @@ +""" +Тесты для 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) + + +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": "привет"}]} diff --git a/to_do_front.txt b/to_do_front.txt new file mode 100644 index 0000000..4fe7b74 --- /dev/null +++ b/to_do_front.txt @@ -0,0 +1,834 @@ +# План разработки фронтенда и бэкенда 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:** [x]
+- **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`. |
+
+---
+
+### Шаг 6.7. UX/UI полировка всех оставшихся пользовательских страниц
+**Цель:** Единообразие, состояния loading, обработка ошибок, фильтры, empty state, мобильная адаптивность на всех страницах пользовательского кабинета и авторизации.
+
+| Файл | Действие |
+|------|----------|
+| `templates/user/dashboard.html` | **Доработать.** Добавить автообновление виджетов квоты и статистики каждые 60 сек. Убедиться, что skeleton присутствует на всех виджетах при первой загрузке. Добавить retry-кнопку на виджет «Последние сессии» при ошибке. Empty state — ссылка на `/asr`. |
+| `templates/user/asr.html` | **Доработать.** Добавить валидацию размера файла перед отправкой (лимит из `Plan.max_audio_duration_sec` или hardcoded). Кнопки «Копировать» и «Скачать .txt» в блоках результата — добавить `UI.setLoading` на время операции. WS-индикатор: добавить пульсацию точки при статусе `connecting`. Блок «История рядом» — добавить skeleton и empty state (сейчас только skeleton). |
+| `templates/user/history.html` | **Доработать.** Добавить фильтр по дате (с/по — input type="date"). Добавить поиск по ID сессии (input с `debounce` 300мс). Убедиться, что hover-эффект на строках таблицы работает через CSS. Пагинация «Загрузить ещё» — `UI.setLoading` на кнопку. |
+| `templates/user/history_detail.html` | **Доработать.** Кнопки «Скачать JSON/TXT» — добавить `UI.setLoading` во время генерации Blob. Pretty-print JSON с CSS-подсветкой синтаксиса (простая, через span color). Кнопка «Назад к истории» — стилизовать как `btn-ghost`. |
+| `templates/user/subscription.html` | **Доработать.** Карточки доступных тарифов — добавить skeleton при загрузке, выделение текущего тарифа рамкой. История платежей — таблица с badge-статусами (pending/succeeded/cancelled), empty state. Toggle auto-renew — `UI.setLoading` + `UI.toast`. |
+| `templates/user/profile.html` | **Доработать.** Кнопка «Сохранить» — `UI.setLoading`. Удаление аккаунта — `UI.confirmDialog` с требованием ввода слова «УДАЛИТЬ» в поле подтверждения. Inline-ошибки под полями (сейчас только для fullName). Пароль: проверка min 8 символов, совпадение new/confirm. |
+| `templates/user/api_keys.html` | **Доработать.** Модалка создания — `UI.setLoading` на кнопке «Создать», валидация имени (не пустое). Показ ключа один раз — кнопка «Копировать» с обратной связью (`UI.toast('Ключ скопирован')`). Отзыв — `UI.confirmDialog`. Список — skeleton + empty state. |
+| `templates/auth/login.html` | **Доработать.** Кнопка «Войти» — `UI.setLoading`. Ошибки API — `UI.toast(msg, 'error')` под формой (сейчас `UI.toast`). Автофокус на поле email при загрузке. |
+| `templates/auth/register.html` | **Доработать.** Аналогично login: `UI.setLoading`, `UI.toast`, inline-проверка паролей (min 8), автофокус на email. |
+| `templates/user/base_user.html` | **Доработать.** Мобильная адаптивность: проверить, что на < 768px sidebar скрыт, bottom tab-bar виден, контент не обрезается. Desktop — sidebar 240px фиксированно. Подсветка активного пункта без ошибок префиксов URL. |
+
+**Общие правила для всех пользовательских страниц (чек-лист):**
+- [x] Все кнопки, вызывающие `fetch`/`apiFetch`, имеют состояние loading (`.spinner` + `disabled`).
+- [x] Все опасные действия (отзыв ключа, отмена подписки, удаление аккаунта) требуют `UI.confirmDialog` (никаких `window.confirm`).
+- [x] Все ошибки API отображаются через `UI.toast(msg, 'error')`.
+- [x] Все успешные действия — `UI.toast(msg, 'success')`.
+- [x] Все списки/таблицы имеют skeleton при загрузке и empty state при отсутствии данных.
+- [x] Все таблицы имеют hover-эффект на строках (`tr:hover` — фон чуть светлее).
+- [x] Фильтры и поиск используют debounce (минимум 300мс).
+- [x] Пагинация: кнопка «Загрузить ещё» (не номера страниц).
+- [x] Мобильная адаптивность: контент не выходит за границы экрана 375px–1440px.
+- [x] Access token синхронизирован с cookie (`auth.js`) для корректной работы SSR-переходов.
+
+**Результат шага 6.7:** выполнена UX/UI полировка всех пользовательских страниц. Добавлены: автообновление dashboard каждые 60 сек, валидация размера файла (max 100 МБ), пульсация WS-индикатора при подключении, setLoading на кнопки копирования/скачивания, фильтры по дате и поиск по ID с debounce в истории, CSS-подсветка синтаксиса JSON в деталях сессии, skeleton и empty state для тарифов/платежей, confirmDialog с обязательным вводом «УДАЛИТЬ» для удаления аккаунта, минимум 8 символов для пароля, автофокус на поле email при загрузке login/register, корректная подсветка активного пункта меню для вложенных URL (history/*).
+
+**Дополнительно выполненные мероприятия (вне плана 6.7):**
+
+| # | Задача | Файлы | Результат |
+|---|--------|-------|-----------|
+| 6.7.доп.1 | **Исправление мобильной адаптивности.** Исправлено поведение sidebar на `base_user.html`: на экранах <768px sidebar скрывается, появляется bottom tab-bar, основной контент не пропадает. | `templates/user/base_user.html` | Мобильная версия работает корректно. |
+| 6.7.доп.2 | **Исправление структуры скриптов.** Исправлены синтаксические ошибки (пропущенные закрывающие теги ``) в `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 — `