diff --git a/backend/app/services/note.py b/backend/app/services/note.py index ebbe83a..f493c91 100644 --- a/backend/app/services/note.py +++ b/backend/app/services/note.py @@ -262,7 +262,10 @@ class NoteGenerator: raise Exception(f"不支持的转写器:{self.transcriber_type}") logger.info(f"使用转写器:{self.transcriber_type}") - return get_transcriber(transcriber_type=self.transcriber_type) + return get_transcriber( + transcriber_type=self.transcriber_type, + model_size=self.model_size, + ) def _get_gpt(self, model_name: Optional[str], provider_id: Optional[str]) -> GPT: """ diff --git a/backend/app/transcriber/transcriber_provider.py b/backend/app/transcriber/transcriber_provider.py index 0440bc8..f4a7b50 100644 --- a/backend/app/transcriber/transcriber_provider.py +++ b/backend/app/transcriber/transcriber_provider.py @@ -1,5 +1,6 @@ import os import platform +import threading from enum import Enum from app.transcriber.groq import GroqTranscriber @@ -38,17 +39,29 @@ _transcribers = { TranscriberType.GROQ: None, } +# Cache instances together with their constructor configuration. The +# transcriber choice and Whisper model size can be changed from the frontend, +# so caching by transcriber type alone would keep using the first loaded model. +_transcriber_configs = {key: None for key in _transcribers} +_transcriber_init_lock = threading.Lock() + # 公共实例初始化函数 def _init_transcriber(key: TranscriberType, cls, *args, **kwargs): - if _transcribers[key] is None: - logger.info(f'创建 {cls.__name__} 实例: {key}') - try: - _transcribers[key] = cls(*args, **kwargs) + init_config = (args, tuple(sorted(kwargs.items()))) + with _transcriber_init_lock: + instance = _transcribers[key] + if instance is None or _transcriber_configs[key] != init_config: + action = "创建" if instance is None else "按新配置重新创建" + logger.info(f'{action} {cls.__name__} 实例: {key}') + try: + new_instance = cls(*args, **kwargs) + except Exception as e: + logger.error(f"{cls.__name__} 创建失败: {e}") + raise + _transcribers[key] = new_instance + _transcriber_configs[key] = init_config logger.info(f'{cls.__name__} 创建成功') - except Exception as e: - logger.error(f"{cls.__name__} 创建失败: {e}") - raise - return _transcribers[key] + return _transcribers[key] # 各类型获取方法 def get_groq_transcriber(): diff --git a/backend/tests/test_whisper_config_forwarding.py b/backend/tests/test_whisper_config_forwarding.py new file mode 100644 index 0000000..7d2ec81 --- /dev/null +++ b/backend/tests/test_whisper_config_forwarding.py @@ -0,0 +1,82 @@ +import os +import pathlib +import subprocess +import sys +import textwrap + + +ROOT = pathlib.Path(__file__).resolve().parents[1] + + +def _run_isolated(script: str) -> None: + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join( + filter(None, [str(ROOT), env.get("PYTHONPATH")]) + ) + result = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + cwd=ROOT, + env=env, + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +def test_note_generator_forwards_configured_whisper_model_size(): + _run_isolated( + """ + from app.services import note as note_service + + calls = [] + + def fake_get_transcriber(**kwargs): + calls.append(kwargs) + return object() + + note_service.get_transcriber = fake_get_transcriber + + generator = note_service.NoteGenerator.__new__(note_service.NoteGenerator) + generator.transcriber_type = "fast-whisper" + generator.model_size = "large-v3-turbo" + + generator._init_transcriber() + + assert calls == [ + { + "transcriber_type": "fast-whisper", + "model_size": "large-v3-turbo", + } + ] + """ + ) + + +def test_whisper_cache_is_rebuilt_when_model_size_changes(): + _run_isolated( + """ + from app.transcriber import transcriber_provider as provider + + class FakeWhisperTranscriber: + def __init__(self, model_size, device): + self.model_size = model_size + self.device = device + + provider._transcribers = {key: None for key in provider._transcribers} + provider._transcriber_configs = { + key: None for key in provider._transcribers + } + provider.WhisperTranscriber = FakeWhisperTranscriber + + base = provider.get_whisper_transcriber("base", device="cpu") + turbo = provider.get_whisper_transcriber("large-v3-turbo", device="cpu") + turbo_again = provider.get_whisper_transcriber( + "large-v3-turbo", device="cpu" + ) + + assert base.model_size == "base" + assert turbo.model_size == "large-v3-turbo" + assert turbo is not base + assert turbo_again is turbo + """ + )