@@ -250,9 +250,11 @@ class DummyEngine:
250250 def __init__ (self ):
251251 self ._engine = types .SimpleNamespace (tokenizer = DummyTokenizer ())
252252 self .input_lengths = []
253+ self .generate_kwargs = []
253254
254255 def generate (self , inputs , ** kwargs ):
255256 self .input_lengths .extend (len (audio ) for audio in inputs )
257+ self .generate_kwargs .append (kwargs )
256258 return [{"text" : "hello" }]
257259
258260
@@ -318,6 +320,21 @@ def test_client_endpoint_mode_emits_partial_at_first_decode_threshold():
318320 assert session .vllm_engine .input_lengths == [int (0.48 * session .sample_rate )]
319321
320322
323+ def test_postprocess_hotwords_correct_final_text_without_model_hotword_bias ():
324+ module = load_service_module ()
325+ session = make_client_endpoint_session (module )
326+ session .asr_kwargs = {
327+ "postprocess_hotwords" : {"hello" : "Tool" },
328+ "return_postprocess_hotword_matches" : True ,
329+ }
330+
331+ session .add_audio (np .zeros (int (0.4 * session .sample_rate ), dtype = np .int16 ).tobytes ())
332+ result = session .commit ()
333+
334+ assert result ["sentences" ] == [{"text" : "Tool" , "start" : 0 , "end" : 400 }]
335+ assert session .vllm_engine .generate_kwargs [- 1 ].get ("hotwords" ) is None
336+
337+
321338def test_client_commits_short_utterances_once_with_monotonic_timestamps ():
322339 module = load_service_module ()
323340 session = make_client_endpoint_session (module )
@@ -447,6 +464,121 @@ async def send(self, message):
447464 assert received_audio == [b"first" , b"second" ]
448465
449466
467+ def test_handler_accepts_postprocess_hotwords_without_model_hotword_bias (monkeypatch ):
468+ module = load_service_module ()
469+ session_kwargs = []
470+
471+ class ProtocolSession :
472+ def __init__ (self , vllm_engine , asr_kwargs , * args , ** kwargs ):
473+ self .is_active = False
474+ self .asr_kwargs = asr_kwargs
475+ session_kwargs .append (dict (asr_kwargs ))
476+
477+ def reset (self ):
478+ pass
479+
480+ def commit (self ):
481+ return {"is_final" : True , "asr_kwargs" : dict (self .asr_kwargs )}
482+
483+ class FakeWebSocket :
484+ remote_address = ("127.0.0.1" , 12345 )
485+
486+ def __init__ (self ):
487+ self .messages = iter (
488+ [
489+ "START" ,
490+ "POSTPROCESS_HOTWORDS:hello=>Tool,哈囉=>客製化" ,
491+ "COMMIT" ,
492+ ]
493+ )
494+ self .sent = []
495+
496+ def __aiter__ (self ):
497+ return self
498+
499+ async def __anext__ (self ):
500+ try :
501+ return next (self .messages )
502+ except StopIteration as error :
503+ raise StopAsyncIteration from error
504+
505+ async def send (self , message ):
506+ self .sent .append (json .loads (message ))
507+
508+ monkeypatch .setattr (module , "load_models" , lambda args : (object (), {}, None , None ))
509+ monkeypatch .setattr (module , "ClientEndpointVAD" , lambda : object (), raising = False )
510+ monkeypatch .setattr (module , "create_speaker_tracker" , lambda model , args : None )
511+ monkeypatch .setattr (module , "RealtimeASRSession" , ProtocolSession )
512+ websocket = FakeWebSocket ()
513+ args = types .SimpleNamespace (
514+ decode_interval = 0.48 ,
515+ partial_window_sec = 15.0 ,
516+ endpoint_mode = "client" ,
517+ log_session_stats_interval = 0.0 ,
518+ )
519+
520+ asyncio .run (module .handle_client (websocket , args ))
521+
522+ assert session_kwargs == [{}]
523+ assert websocket .sent [1 ] == {
524+ "event" : "postprocess_hotwords_set" ,
525+ "postprocess_hotwords" : {"hello" : "Tool" , "哈囉" : "客製化" },
526+ }
527+ final = [message for message in websocket .sent if message .get ("is_final" )][0 ]
528+ assert final ["asr_kwargs" ] == {
529+ "postprocess_hotwords" : {"hello" : "Tool" , "哈囉" : "客製化" },
530+ "postprocess_hotword_fuzzy" : False ,
531+ "return_postprocess_hotword_matches" : True ,
532+ }
533+ assert "hotwords" not in final ["asr_kwargs" ]
534+
535+
536+ def test_load_models_reads_postprocess_hotword_file (monkeypatch , tmp_path ):
537+ module = load_service_module ()
538+ module ._vllm_engine = None
539+ module ._asr_kwargs = None
540+ module ._vad_model = None
541+ module ._spk_model = None
542+ hotword_file = tmp_path / "postprocess_hotwords.txt"
543+ hotword_file .write_text ("hello=>Tool\n 哈囉=>客製化\n " , encoding = "utf-8" )
544+ auto_model_calls = []
545+
546+ import funasr
547+
548+ monkeypatch .setattr (
549+ funasr ,
550+ "AutoModel" ,
551+ lambda model , ** kwargs : auto_model_calls .append (model ) or object (),
552+ )
553+ vllm_stub = types .ModuleType ("funasr.auto.auto_model_vllm" )
554+ vllm_stub .AutoModelVLLM = lambda ** kwargs : object ()
555+ monkeypatch .setitem (sys .modules , "funasr.auto.auto_model_vllm" , vllm_stub )
556+
557+ args = types .SimpleNamespace (
558+ model = "FunAudioLLM/Fun-ASR-Nano-2512" ,
559+ hub = "ms" ,
560+ device = "cpu" ,
561+ dtype = "fp32" ,
562+ tensor_parallel_size = 1 ,
563+ gpu_memory_utilization = 0.8 ,
564+ max_model_len = 2048 ,
565+ hotword_file = "" ,
566+ postprocess_hotword_file = str (hotword_file ),
567+ language = None ,
568+ enable_spk = False ,
569+ endpoint_mode = "client" ,
570+ )
571+
572+ _ , asr_kwargs , _ , _ = module .load_models (args )
573+
574+ assert asr_kwargs == {
575+ "postprocess_hotword_file" : str (hotword_file ),
576+ "postprocess_hotword_fuzzy" : False ,
577+ "return_postprocess_hotword_matches" : True ,
578+ }
579+ assert auto_model_calls == []
580+
581+
450582def test_two_hour_session_keeps_audio_bounded_and_duration_absolute ():
451583 module = load_service_module ()
452584 sample_rate = 10
0 commit comments