Skip to content

Commit ff70f61

Browse files
authored
Merge pull request #45 from quanru/feat/recognition-provider-foundation
refactor(asr): add recognition provider registry
2 parents 3e78cea + 2c291b7 commit ff70f61

6 files changed

Lines changed: 171 additions & 35 deletions

File tree

‎src/doubao_input/app.py‎

Lines changed: 25 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,6 @@
2626
from doubao_input.result import RecentResult
2727
from doubao_input.doubao.params_store import ParamsStore
2828
from doubao_input.doubao.transcription import TranscriptionManager
29-
from doubao_input.doubao.asr_client import ASRClient
3029
from doubao_input.doubao.volcengine_asr_client import VolcengineASRClient
3130
from doubao_input.doubao.volcengine_credentials import (
3231
VolcengineCredentials, VolcengineCredentialsStore,
@@ -48,6 +47,7 @@
4847
from doubao_input.polish_preview import PolishPreview
4948
from doubao_input.trigger.escape_guard import EscapeGuard
5049
from doubao_input.diagnostics import DiagnosticTrace, report as diagnostic_report
50+
from doubao_input.recognition_providers import recognition_provider
5151

5252
logger = logging.getLogger(__name__)
5353

@@ -347,29 +347,29 @@ def apply_runtime(value, strict):
347347
self._control.refresh()
348348

349349
def _new_transcription_manager(self):
350-
if self.settings.asr_provider == "volcengine":
351-
return TranscriptionManager(self.app_state,
352-
asr_client=VolcengineASRClient(),
353-
credential_store=VolcengineCredentialsStore,
354-
interactive_auth=False, clear_rejected_credentials=False)
355-
return TranscriptionManager(self.app_state, asr_client=ASRClient(),
356-
credential_store=ParamsStore, interactive_auth=True,
357-
clear_rejected_credentials=True)
350+
provider = recognition_provider(self.settings.asr_provider)
351+
return TranscriptionManager(
352+
self.app_state,
353+
asr_client=provider.new_client(),
354+
credential_store=provider.credential_store,
355+
interactive_auth=provider.interactive_auth,
356+
clear_rejected_credentials=provider.clear_rejected_credentials,
357+
)
358358

359359
def _configure_recognition_backend(self):
360-
if self.settings.asr_provider == "volcengine":
361-
self._tm.configure_backend(VolcengineASRClient(),
362-
VolcengineCredentialsStore, interactive_auth=False,
363-
clear_rejected_credentials=False)
364-
else:
365-
self._tm.configure_backend(ASRClient(), ParamsStore,
366-
interactive_auth=True, clear_rejected_credentials=True)
360+
provider = recognition_provider(self.settings.asr_provider)
361+
self._tm.configure_backend(
362+
provider.new_client(),
363+
provider.credential_store,
364+
interactive_auth=provider.interactive_auth,
365+
clear_rejected_credentials=provider.clear_rejected_credentials,
366+
)
367367

368368
def _recognition_ready(self):
369369
try:
370-
store = (VolcengineCredentialsStore if self.settings.asr_provider == "volcengine"
371-
else ParamsStore)
372-
return store.has_saved()
370+
return recognition_provider(
371+
self.settings.asr_provider
372+
).credential_store.has_saved()
373373
except (OSError, ValueError):
374374
return False
375375

@@ -765,10 +765,8 @@ def _summary(self):
765765
"microphone_id": self.settings.microphone,
766766
"microphone_ok": bool(setup and setup.microphone_ok),
767767
"asr_provider": self.settings.asr_provider,
768-
"asr_provider_name": tr(
769-
"Volcengine official API", "火山引擎官方 API")
770-
if self.settings.asr_provider == "volcengine" else
771-
tr("Doubao account", "豆包账号"),
768+
"asr_provider_name": recognition_provider(
769+
self.settings.asr_provider).name,
772770
"voice_test_ok": bool(setup and setup.voice_ok),
773771
"onboarding_complete": self.settings.onboarding_complete,
774772
"result": self.recent.text, "status": self.recent.status}
@@ -816,9 +814,10 @@ def _sign_out(self):
816814
if self._login_window:
817815
self._login_window.destroy()
818816
self._login_window = None
819-
official = getattr(getattr(self, "settings", None), "asr_provider", "doubao") == "volcengine"
820-
store = (VolcengineCredentialsStore if official
821-
else ParamsStore)
817+
provider = recognition_provider(
818+
getattr(getattr(self, "settings", None), "asr_provider", "doubao"))
819+
official = provider.id == "volcengine"
820+
store = provider.credential_store
822821
try:
823822
store.clear()
824823
except OSError as error:
Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
"""Recognition provider registry shared by settings, UI, and runtime wiring."""
2+
3+
from __future__ import annotations
4+
5+
from dataclasses import dataclass
6+
from typing import Callable, Protocol
7+
8+
from doubao_input.i18n import tr
9+
10+
11+
class CredentialStore(Protocol):
12+
@classmethod
13+
def load(cls): ...
14+
15+
@classmethod
16+
def clear(cls) -> None: ...
17+
18+
@classmethod
19+
def has_saved(cls) -> bool: ...
20+
21+
22+
@dataclass(frozen=True)
23+
class RecognitionProvider:
24+
id: str
25+
name_en: str
26+
name_zh: str
27+
client_factory: Callable[[], object]
28+
credential_store_factory: Callable[[], type[CredentialStore]]
29+
interactive_auth: bool
30+
clear_rejected_credentials: bool
31+
32+
@property
33+
def name(self) -> str:
34+
return tr(self.name_en, self.name_zh)
35+
36+
def new_client(self):
37+
return self.client_factory()
38+
39+
@property
40+
def credential_store(self) -> type[CredentialStore]:
41+
return self.credential_store_factory()
42+
43+
44+
def _doubao_client():
45+
from doubao_input.doubao.asr_client import ASRClient
46+
return ASRClient()
47+
48+
49+
def _doubao_store():
50+
from doubao_input.doubao.params_store import ParamsStore
51+
return ParamsStore
52+
53+
54+
def _volcengine_client():
55+
from doubao_input.doubao.volcengine_asr_client import VolcengineASRClient
56+
return VolcengineASRClient()
57+
58+
59+
def _volcengine_store():
60+
from doubao_input.doubao.volcengine_credentials import VolcengineCredentialsStore
61+
return VolcengineCredentialsStore
62+
63+
64+
_PROVIDERS = (
65+
RecognitionProvider(
66+
id="doubao",
67+
name_en="Doubao account",
68+
name_zh="豆包账号",
69+
client_factory=_doubao_client,
70+
credential_store_factory=_doubao_store,
71+
interactive_auth=True,
72+
clear_rejected_credentials=True,
73+
),
74+
RecognitionProvider(
75+
id="volcengine",
76+
name_en="Volcengine official API",
77+
name_zh="火山引擎官方 API",
78+
client_factory=_volcengine_client,
79+
credential_store_factory=_volcengine_store,
80+
interactive_auth=False,
81+
clear_rejected_credentials=False,
82+
),
83+
)
84+
85+
RECOGNITION_PROVIDER_IDS = tuple(provider.id for provider in _PROVIDERS)
86+
_PROVIDERS_BY_ID = {provider.id: provider for provider in _PROVIDERS}
87+
88+
89+
def recognition_providers() -> tuple[RecognitionProvider, ...]:
90+
return _PROVIDERS
91+
92+
93+
def recognition_provider(provider_id: str) -> RecognitionProvider:
94+
try:
95+
return _PROVIDERS_BY_ID[provider_id]
96+
except KeyError as error:
97+
raise ValueError(tr(
98+
"Unsupported recognition service",
99+
"不支持的语音识别服务",
100+
)) from error

‎src/doubao_input/settings.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,11 +7,12 @@
77
import tempfile
88
from urllib.parse import urlsplit
99
from doubao_input.i18n import LANGUAGES, tr
10+
from doubao_input.recognition_providers import RECOGNITION_PROVIDER_IDS
1011

1112
KEY_CHOICES = {"Disabled": 0, "Fn": 464, "Ctrl": 29, "Shift": 42,
1213
"Alt": 56, "Meta": 125, "F8": 66, "F9": 67}
1314
REMOVED_SETTING_FIELDS = frozenset({"polish_undo_modifier", "polish_prompt"})
14-
ASR_PROVIDERS = ("doubao", "volcengine")
15+
ASR_PROVIDERS = RECOGNITION_PROVIDER_IDS
1516
CAPTURABLE_KEY_CODES = frozenset(range(2, 249)) | {464}
1617
MODIFIER_KEY_CODES = frozenset({29, 42, 54, 56, 97, 100, 125, 126})
1718
EQUIVALENT_KEY_GROUPS = (

‎src/doubao_input/ui/control_window.py‎

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from doubao_input.ui.style import apply_window_style
1111
from doubao_input.product import VERSION
1212
from doubao_input.settings import ASR_PROVIDERS
13+
from doubao_input.recognition_providers import recognition_providers
1314

1415

1516
class ControlWindow:
@@ -324,10 +325,8 @@ def button(box, text, callback, suggested=False):
324325
provider_row.append(Gtk.Label(
325326
label=tr("Recognition service", "语音识别服务"),
326327
xalign=0, hexpand=True, wrap=True))
327-
self._asr_provider = Gtk.DropDown.new_from_strings([
328-
tr("Doubao account", "豆包账号"),
329-
tr("Volcengine API", "火山引擎 API"),
330-
])
328+
self._asr_provider = Gtk.DropDown.new_from_strings(
329+
[provider.name for provider in recognition_providers()])
331330
self._asr_provider.set_selected(ASR_PROVIDERS.index(
332331
self._actions.summary().get("asr_provider", "doubao")))
333332
self._asr_provider.connect("notify::selected", self._asr_provider_changed)

‎src/doubao_input/ui/settings_window.py‎

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from doubao_input.settings import INPUT_METHODS, WAVEFORM_STYLES
1010
from doubao_input.ui.style import apply_window_style
1111
from doubao_input.settings import ASR_PROVIDERS
12+
from doubao_input.recognition_providers import recognition_providers
1213

1314

1415
class SettingsWindow:
@@ -103,10 +104,8 @@ def run(*_):
103104
"更改会自动保存;选择“系统默认”可跟随桌面的输入设备。")))
104105

105106
section(tr("Recognition service", "语音识别服务"))
106-
self.asr_provider = Gtk.DropDown.new_from_strings([
107-
tr("Doubao account", "豆包账号"),
108-
tr("Volcengine official API", "火山引擎官方 API"),
109-
])
107+
self.asr_provider = Gtk.DropDown.new_from_strings(
108+
[provider.name for provider in recognition_providers()])
110109
self.asr_provider.set_selected(ASR_PROVIDERS.index(settings.asr_provider))
111110
row(tr("Service", "服务"), self.asr_provider)
112111
self.asr_details = Gtk.Box(orientation=Gtk.Orientation.VERTICAL, spacing=10)
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
from unittest import TestCase
2+
3+
from doubao_input.doubao.asr_client import ASRClient
4+
from doubao_input.doubao.params_store import ParamsStore
5+
from doubao_input.doubao.volcengine_asr_client import VolcengineASRClient
6+
from doubao_input.doubao.volcengine_credentials import VolcengineCredentialsStore
7+
from doubao_input.recognition_providers import (
8+
RECOGNITION_PROVIDER_IDS,
9+
recognition_provider,
10+
recognition_providers,
11+
)
12+
13+
14+
class RecognitionProviderRegistryTest(TestCase):
15+
def test_registry_order_drives_persisted_ids_and_ui_names(self):
16+
providers = recognition_providers()
17+
self.assertEqual(
18+
RECOGNITION_PROVIDER_IDS,
19+
tuple(provider.id for provider in providers),
20+
)
21+
self.assertEqual(
22+
[provider.name for provider in providers],
23+
["Doubao account", "Volcengine official API"],
24+
)
25+
26+
def test_registry_owns_runtime_backend_configuration(self):
27+
doubao = recognition_provider("doubao")
28+
self.assertIsInstance(doubao.new_client(), ASRClient)
29+
self.assertIs(doubao.credential_store, ParamsStore)
30+
self.assertTrue(doubao.interactive_auth)
31+
official = recognition_provider("volcengine")
32+
self.assertIsInstance(official.new_client(), VolcengineASRClient)
33+
self.assertIs(official.credential_store, VolcengineCredentialsStore)
34+
self.assertFalse(official.interactive_auth)
35+
36+
def test_unknown_provider_is_rejected(self):
37+
with self.assertRaisesRegex(ValueError, "Unsupported recognition service"):
38+
recognition_provider("missing")

0 commit comments

Comments
 (0)