main
py 193 lines 5.22 KB
Raw
1 from __future__ import annotations
2
3 import asyncio
4 import base64
5 import os
6 import tempfile
7 import warnings
8 from typing import Any
9
10 import whisper
11
12 from helpers import files, plugins
13 from helpers.notification import (
14 NotificationManager,
15 NotificationPriority,
16 NotificationType,
17 )
18 from helpers.print_style import PrintStyle
19 from plugins._whisper_stt.helpers import migration
20
21
22 warnings.filterwarnings("ignore", category=FutureWarning)
23
24
25 PLUGIN_NAME = "_whisper_stt"
26 DEFAULT_CONFIG = {
27 "model_size": "base",
28 "language": "en",
29 "message_mode": "send",
30 "silence_threshold": 0.3,
31 "silence_duration": 1000,
32 "waiting_timeout": 2000,
33 }
34 VALID_MODEL_SIZES = {"tiny", "base", "small", "medium", "large", "turbo"}
35 VALID_MESSAGE_MODES = {"send", "draft"}
36
37 _model = None
38 _model_name = ""
39 is_updating_model = False
40
41
42 def normalize_config(config: dict[str, Any] | None) -> dict[str, Any]:
43 normalized = dict(DEFAULT_CONFIG)
44 if not isinstance(config, dict):
45 return normalized
46
47 model_size = str(config.get("model_size", normalized["model_size"]) or "").strip()
48 if model_size in VALID_MODEL_SIZES:
49 normalized["model_size"] = model_size
50
51 language = str(config.get("language", normalized["language"]) or "").strip()
52 if language:
53 normalized["language"] = language
54
55 message_mode = (
56 str(config.get("message_mode", normalized["message_mode"]) or "")
57 .strip()
58 .lower()
59 )
60 if message_mode in VALID_MESSAGE_MODES:
61 normalized["message_mode"] = message_mode
62
63 try:
64 silence_threshold = float(
65 config.get("silence_threshold", normalized["silence_threshold"])
66 )
67 normalized["silence_threshold"] = min(max(silence_threshold, 0.0), 1.0)
68 except (TypeError, ValueError):
69 pass
70
71 try:
72 silence_duration = int(
73 config.get("silence_duration", normalized["silence_duration"])
74 )
75 if silence_duration > 0:
76 normalized["silence_duration"] = silence_duration
77 except (TypeError, ValueError):
78 pass
79
80 try:
81 waiting_timeout = int(config.get("waiting_timeout", normalized["waiting_timeout"]))
82 if waiting_timeout > 0:
83 normalized["waiting_timeout"] = waiting_timeout
84 except (TypeError, ValueError):
85 pass
86
87 return normalized
88
89
90 def get_config() -> dict[str, Any]:
91 migration.ensure_config_seeded()
92 config = plugins.get_plugin_config(PLUGIN_NAME) or {}
93 return normalize_config(config)
94
95
96 def get_loaded_model_name() -> str:
97 return _model_name
98
99
100 def is_globally_enabled() -> bool:
101 return plugins.determined_toggle_from_paths(
102 True, reversed(plugins.get_plugin_roots(PLUGIN_NAME))
103 )
104
105
106 async def preload(model_name: str | None = None):
107 cfg = get_config()
108 resolved_model = str(model_name or cfg["model_size"])
109 return await _preload(resolved_model)
110
111
112 async def _preload(model_name: str):
113 global _model, _model_name, is_updating_model
114
115 while is_updating_model:
116 await asyncio.sleep(0.1)
117
118 try:
119 is_updating_model = True
120 if not _model or _model_name != model_name:
121 NotificationManager.send_notification(
122 NotificationType.INFO,
123 NotificationPriority.NORMAL,
124 "Loading Whisper model...",
125 display_time=99,
126 group="whisper-preload",
127 )
128 PrintStyle.standard(f"Loading Whisper model: {model_name}")
129 _model = whisper.load_model(
130 name=model_name,
131 download_root=files.get_abs_path("/tmp/models/whisper"),
132 )
133 _model_name = model_name
134 NotificationManager.send_notification(
135 NotificationType.INFO,
136 NotificationPriority.NORMAL,
137 "Whisper model loaded.",
138 display_time=2,
139 group="whisper-preload",
140 )
141 finally:
142 is_updating_model = False
143
144
145 async def is_downloading() -> bool:
146 return is_updating_model
147
148
149 async def is_downloaded() -> bool:
150 return _model is not None
151
152
153 async def transcribe(
154 audio_bytes_b64: str, config: dict[str, Any] | None = None
155 ) -> dict[str, Any]:
156 cfg = normalize_config(config or get_config())
157 return await _transcribe(
158 str(cfg["model_size"]),
159 audio_bytes_b64,
160 language=_resolve_language(str(cfg["language"])),
161 )
162
163
164 async def _transcribe(
165 model_name: str, audio_bytes_b64: str, *, language: str | None = None
166 ) -> dict[str, Any]:
167 await _preload(model_name)
168
169 audio_bytes = base64.b64decode(audio_bytes_b64)
170
171 with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as audio_file:
172 audio_file.write(audio_bytes)
173 temp_path = audio_file.name
174
175 try:
176 kwargs: dict[str, Any] = {"fp16": False}
177 if language:
178 kwargs["language"] = language
179
180 result = _model.transcribe(temp_path, **kwargs) # type: ignore[union-attr]
181 return result if isinstance(result, dict) else {}
182 finally:
183 try:
184 os.remove(temp_path)
185 except Exception:
186 pass
187
188
189 def _resolve_language(language: str) -> str | None:
190 value = language.strip().lower()
191 if not value or value == "auto":
192 return None
193 return value