From b2ba0bb1ee5ac1f5ef212d949884c17670bfc730 Mon Sep 17 00:00:00 2001 From: zhifu gao Date: Tue, 23 Jun 2026 05:40:39 +0800 Subject: [PATCH 1/2] feat: add funasr_asr_python local FunASR (SenseVoice) ASR extension --- .../extension/funasr_asr_python/README.md | 37 ++ .../extension/funasr_asr_python/__init__.py | 8 + .../extension/funasr_asr_python/addon.py | 15 + .../extension/funasr_asr_python/config.py | 63 ++++ .../extension/funasr_asr_python/const.py | 9 + .../extension/funasr_asr_python/extension.py | 334 ++++++++++++++++++ .../funasr_asr_python/funasr_client.py | 167 +++++++++ .../extension/funasr_asr_python/manifest.json | 57 +++ .../extension/funasr_asr_python/property.json | 13 + .../funasr_asr_python/reconnect_manager.py | 90 +++++ .../funasr_asr_python/requirements.txt | 4 + .../funasr_asr_python/tests/__init__.py | 4 + .../funasr_asr_python/tests/test_config.py | 45 +++ 13 files changed, 846 insertions(+) create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/README.md create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/__init__.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/addon.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/const.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/property.json create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/reconnect_manager.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/__init__.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_config.py diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/README.md b/ai_agents/agents/ten_packages/extension/funasr_asr_python/README.md new file mode 100644 index 0000000000..f3c40b35c4 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/README.md @@ -0,0 +1,37 @@ +# FunASR ASR Python Extension + +A local (self-hosted) speech-to-text extension for TEN Framework, powered by +[FunASR](https://github.com/modelscope/FunASR) — an open-source speech toolkit +from Tongyi Lab with strong multilingual ASR (Chinese, Cantonese, English, +Japanese, Korean and more). The model runs locally on CPU or CUDA; **no API key +is required** and audio never leaves your machine. + +It mirrors the existing `whisper_stt_python` extension, swapping faster-whisper +for a local FunASR model — a natural fit for Chinese / Asian-language agents. + +## Features + +- Local, self-hosted ASR — no cloud, no API key, no per-minute billing. +- Strong Chinese / multilingual recognition; default model **SenseVoice-Small** + auto-detects the spoken language and emits inverse-text-normalized text. +- Swappable models via config (e.g. the flagship `FunAudioLLM/Fun-ASR-Nano-2512` + on GPU, or `paraformer-zh` for Chinese with timestamps). +- CPU or CUDA via the `device` parameter. + +## Configuration (`property.json` → `params`) + +| Param | Default | Description | +|---|---|---| +| `model` | `iic/SenseVoiceSmall` | FunASR model id. Use `FunAudioLLM/Fun-ASR-Nano-2512` (flagship LLM-ASR) on GPU, or `paraformer-zh` for Chinese. | +| `device` | `cpu` | `cpu` or `cuda`. | +| `language` | `auto` | Language hint, or `auto` to detect. | +| `use_itn` | `true` | Apply inverse text normalization. | +| `sample_rate` | `16000` | Input PCM sample rate (16-bit mono). | + +## Requirements + +``` +pip install funasr +``` + +The model is downloaded automatically on first run. diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/__init__.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/__init__.py new file mode 100644 index 0000000000..68ee8ef51b --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/__init__.py @@ -0,0 +1,8 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +from . import addon + +__all__ = ["addon"] diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/addon.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/addon.py new file mode 100644 index 0000000000..2ab3f62cfd --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/addon.py @@ -0,0 +1,15 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +from ten_runtime import Addon, register_addon_as_extension, TenEnv + + +@register_addon_as_extension("funasr_asr_python") +class FunASRExtensionAddon(Addon): + def on_create_instance(self, ten_env: TenEnv, name: str, context) -> None: + from .extension import FunASRExtension + + ten_env.log_info("FunASRExtensionAddon on_create_instance") + ten_env.on_create_instance_done(FunASRExtension(name), context) diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py new file mode 100644 index 0000000000..4b14bb0de0 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py @@ -0,0 +1,63 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +from typing import Dict, Any +from pydantic import BaseModel, Field +from ten_ai_base.utils import encrypt + + +class FunASRConfig(BaseModel): + """FunASR ASR Configuration""" + + # Debugging and dumping + dump: bool = False + dump_path: str = "/tmp" + + # Finalize mode: "disconnect" or "silence" + finalize_mode: str = "disconnect" + silence_duration_ms: int = 1000 + + # Vendor parameters (pass-through design) + params: Dict[str, Any] = Field(default_factory=dict) + + def update(self, params: Dict[str, Any]) -> None: + """Update configuration with additional parameters.""" + for key, value in params.items(): + if hasattr(self, key): + setattr(self, key, value) + + def to_json(self, sensitive_handling: bool = False) -> str: + """Convert config to JSON string with optional sensitive data handling.""" + config_dict = self.model_dump() + if sensitive_handling and config_dict.get("params"): + # Mask sensitive keys + for key in ["api_key", "key", "token", "secret"]: + if key in config_dict["params"] and config_dict["params"][key]: + config_dict["params"][key] = encrypt( + config_dict["params"][key] + ) + return str(config_dict) + + @property + def normalized_language(self) -> str: + """Convert language code to normalized format""" + language_map = { + "zh": "zh-CN", + "en": "en-US", + "ja": "ja-JP", + "ko": "ko-KR", + "yue": "yue-CN", + "de": "de-DE", + "fr": "fr-FR", + "ru": "ru-RU", + "es": "es-ES", + "pt": "pt-PT", + "it": "it-IT", + "hi": "hi-IN", + "ar": "ar-AE", + } + params_dict = self.params or {} + language_code = params_dict.get("language", "") or "" + return language_map.get(language_code, language_code) diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/const.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/const.py new file mode 100644 index 0000000000..7a81e97b78 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/const.py @@ -0,0 +1,9 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +MODULE_NAME_ASR = "asr" +DUMP_FILE_NAME = "funasr_asr_in.pcm" +LOG_CATEGORY_VENDOR = "vendor" +LOG_CATEGORY_KEY_POINT = "key_point" diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py new file mode 100644 index 0000000000..8d3d4a4594 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py @@ -0,0 +1,334 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +import os +from datetime import datetime +from typing import Optional +from typing_extensions import override + +from ten_runtime import AsyncTenEnv, AudioFrame +from ten_ai_base.asr import ( + ASRBufferConfig, + ASRBufferConfigModeKeep, + ASRResult, + AsyncASRBaseExtension, +) +from ten_ai_base.message import ( + ModuleError, + ModuleErrorVendorInfo, + ModuleErrorCode, +) +from ten_ai_base.const import ( + LOG_CATEGORY_VENDOR, + LOG_CATEGORY_KEY_POINT, +) +from ten_ai_base.dumper import Dumper + +from .config import FunASRConfig +from .const import MODULE_NAME_ASR, DUMP_FILE_NAME +from .reconnect_manager import ReconnectManager +from .funasr_client import FunASRClient + + +class FunASRExtension(AsyncASRBaseExtension): + """FunASR ASR Extension using a local FunASR model (SenseVoice / Fun-ASR-Nano / Paraformer).""" + + def __init__(self, name: str): + super().__init__(name) + self.config: Optional[FunASRConfig] = None + self.client: Optional[FunASRClient] = None + self.audio_dumper: Optional[Dumper] = None + self.reconnect_manager: Optional[ReconnectManager] = None + self.sent_user_audio_duration_ms_before_last_reset: int = 0 + self.last_finalize_timestamp: int = 0 + + @override + def vendor(self) -> str: + """Get ASR vendor name""" + return "funasr" + + @override + async def on_init(self, ten_env: AsyncTenEnv) -> None: + await super().on_init(ten_env) + + # Initialize reconnection manager + self.reconnect_manager = ReconnectManager(logger=ten_env) + + config_json, _ = await ten_env.get_property_to_json("") + + try: + self.config = FunASRConfig.model_validate_json(config_json) + self.config.update(self.config.params) + ten_env.log_info( + f"config: {self.config.to_json(sensitive_handling=True)}", + category=LOG_CATEGORY_KEY_POINT, + ) + + # Initialize audio dumper if enabled + if self.config.dump: + dump_file_path = os.path.join( + self.config.dump_path, DUMP_FILE_NAME + ) + self.audio_dumper = Dumper(dump_file_path) + await self.audio_dumper.start() + + except Exception as e: + ten_env.log_error(f"Invalid FunASR config: {e}") + self.config = FunASRConfig.model_validate_json("{}") + await self.send_asr_error( + ModuleError( + module=MODULE_NAME_ASR, + code=ModuleErrorCode.FATAL_ERROR.value, + message=str(e), + ), + ) + + @override + async def on_deinit(self, ten_env: AsyncTenEnv) -> None: + await super().on_deinit(ten_env) + if self.audio_dumper: + await self.audio_dumper.stop() + self.audio_dumper = None + + @override + async def start_connection(self) -> None: + """Start ASR connection (load the local FunASR model)""" + assert self.config is not None + self.ten_env.log_info("Starting FunASR connection") + + try: + # Stop existing connection + if self.is_connected(): + await self.stop_connection() + + # Get configuration parameters + model = self.config.params.get("model", "iic/SenseVoiceSmall") + device = self.config.params.get("device", "cpu") + language = self.config.params.get("language", "auto") + use_itn = self.config.params.get("use_itn", True) + sample_rate = self.config.params.get("sample_rate", 16000) + + # Create client + self.client = FunASRClient( + model=model, + device=device, + language=language, + use_itn=use_itn, + sample_rate=sample_rate, + on_result_callback=self._on_result, + on_error_callback=self._on_error, + logger=self.ten_env, + ) + + await self.client.connect() + + # Mark connection successful + if self.reconnect_manager: + self.reconnect_manager.mark_connection_successful() + + # Reset timeline + self.sent_user_audio_duration_ms_before_last_reset += ( + self.audio_timeline.get_total_user_audio_duration() + ) + self.audio_timeline.reset() + + self.ten_env.log_info( + "FunASR connection established", + category=LOG_CATEGORY_VENDOR, + ) + + except Exception as e: + self.ten_env.log_error(f"Failed to start FunASR connection: {e}") + await self.send_asr_error( + ModuleError( + module=MODULE_NAME_ASR, + code=ModuleErrorCode.FATAL_ERROR.value, + message=str(e), + ), + ) + + @override + async def stop_connection(self) -> None: + """Stop ASR connection""" + self.ten_env.log_info("Stopping FunASR connection") + try: + if self.client: + await self.client.disconnect() + self.client = None + self.ten_env.log_info("FunASR connection stopped") + except Exception as e: + self.ten_env.log_error(f"Error stopping FunASR connection: {e}") + + @override + def is_connected(self) -> bool: + """Check connection status""" + return self.client is not None and self.client.is_connected() + + @override + def input_audio_sample_rate(self) -> int: + """Input audio sample rate""" + assert self.config is not None + return self.config.params.get("sample_rate", 16000) + + @override + def buffer_strategy(self) -> ASRBufferConfig: + """Buffer strategy configuration""" + return ASRBufferConfigModeKeep(byte_limit=1024 * 1024 * 10) + + @override + async def send_audio( + self, frame: AudioFrame, _session_id: str | None + ) -> bool: + """Send audio data""" + if not self.is_connected() or not self.client: + return False + + try: + buf = frame.lock_buf() + audio_data = bytes(buf) + + # Dump audio data + if self.audio_dumper: + await self.audio_dumper.push_bytes(audio_data) + + # Send to client + await self.client.send_audio(audio_data) + + frame.unlock_buf(buf) + return True + + except Exception as e: + self.ten_env.log_error(f"Error sending audio to FunASR: {e}") + frame.unlock_buf(buf) + return False + + @override + async def finalize(self, _session_id: str | None) -> None: + """Finalize recognition""" + assert self.config is not None + + self.last_finalize_timestamp = int(datetime.now().timestamp() * 1000) + self.ten_env.log_debug( + f"FunASR finalize start at {self.last_finalize_timestamp}" + ) + + finalize_mode = self.config.finalize_mode + if finalize_mode == "disconnect": + await self._handle_finalize_disconnect() + elif finalize_mode == "silence": + await self._handle_finalize_silence() + else: + raise ValueError(f"invalid finalize mode: {finalize_mode}") + + async def _handle_finalize_disconnect(self) -> None: + """Handle disconnect mode finalization""" + if self.client: + await self.client.finalize() + self.ten_env.log_debug( + "FunASR finalize completed (disconnect mode)" + ) + + async def _handle_finalize_silence(self) -> None: + """Handle silence mode finalization""" + if self.client and self.config: + # Process any remaining audio + await self.client.finalize() + self.ten_env.log_debug("FunASR finalize completed (silence mode)") + + async def _handle_reconnect(self) -> None: + """Handle reconnection""" + if not self.reconnect_manager: + self.ten_env.log_error("ReconnectManager not initialized") + return + + success = await self.reconnect_manager.handle_reconnect( + connection_func=self.start_connection, + error_handler=self.send_asr_error, + ) + + if success: + self.ten_env.log_debug( + "Reconnection attempt initiated successfully" + ) + else: + info = self.reconnect_manager.get_attempts_info() + self.ten_env.log_debug( + f"Reconnection attempt failed. Status: {info}" + ) + + async def _finalize_end(self) -> None: + """Handle finalization end logic""" + if self.last_finalize_timestamp != 0: + timestamp = int(datetime.now().timestamp() * 1000) + latency = timestamp - self.last_finalize_timestamp + self.ten_env.log_debug( + f"FunASR finalize end at {timestamp}, latency: {latency}ms" + ) + self.last_finalize_timestamp = 0 + await self.send_asr_finalize_end() + + async def _on_result( + self, + text: str, + start_ms: int, + duration_ms: int, + language: str, + final: bool, + ) -> None: + """Handle recognition result callback""" + try: + if not text: + return + + # Calculate actual start time using audio timeline + actual_start_ms = int( + self.audio_timeline.get_audio_duration_before_time(start_ms) + + self.sent_user_audio_duration_ms_before_last_reset + ) + + # Create ASR result + asr_result = ASRResult( + text=text, + final=final, + start_ms=actual_start_ms, + duration_ms=duration_ms, + language=( + self.config.normalized_language if self.config else language + ), + words=[], + ) + + await self.send_asr_result(asr_result) + + # Handle finalize end + if final: + await self._finalize_end() + + except Exception as e: + self.ten_env.log_error(f"Error processing FunASR result: {e}") + + async def _on_error(self, error_msg: str) -> None: + """Handle error callback""" + self.ten_env.log_error( + f"vendor_error: {error_msg}", + category=LOG_CATEGORY_VENDOR, + ) + + # Send error information + await self.send_asr_error( + ModuleError( + module=MODULE_NAME_ASR, + code=ModuleErrorCode.NON_FATAL_ERROR.value, + message=error_msg, + ), + ModuleErrorVendorInfo( + vendor=self.vendor(), + code="unknown", + message=error_msg, + ), + ) + + # Attempt reconnection + await self._handle_reconnect() diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py new file mode 100644 index 0000000000..e9bd769a9e --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py @@ -0,0 +1,167 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +import asyncio +import numpy as np +from typing import Optional, Callable +from funasr import AutoModel +from funasr.utils.postprocess_utils import rich_transcription_postprocess + + +class FunASRClient: + """Client for local FunASR ASR processing (SenseVoice / Fun-ASR-Nano / Paraformer).""" + + def __init__( + self, + model: str = "iic/SenseVoiceSmall", + device: str = "cpu", + language: str = "auto", + use_itn: bool = True, + sample_rate: int = 16000, + on_result_callback: Optional[Callable] = None, + on_error_callback: Optional[Callable] = None, + logger: Optional[any] = None, + ): + self.model_id = model + self.device = device + self.language = language + self.use_itn = use_itn + self.sample_rate = sample_rate + self.on_result_callback = on_result_callback + self.on_error_callback = on_error_callback + self.logger = logger + + self.model: Optional[AutoModel] = None + self.audio_buffer = bytearray() + self.is_connected_flag = False + self.processing_lock = asyncio.Lock() + + # Buffer settings + self.min_audio_length_ms = 1000 # Minimum 1 second of audio + self.max_audio_length_ms = 30000 # Maximum 30 seconds per chunk + + async def connect(self) -> None: + """Load the local FunASR model (off the event loop).""" + try: + if self.logger: + self.logger.log_info( + f"Loading FunASR model: {self.model_id} on {self.device}" + ) + + loop = asyncio.get_event_loop() + self.model = await loop.run_in_executor( + None, + lambda: AutoModel( + model=self.model_id, + device=self.device, + disable_update=True, + ), + ) + + self.is_connected_flag = True + if self.logger: + self.logger.log_info("FunASR model loaded successfully") + + except Exception as e: + self.is_connected_flag = False + if self.logger: + self.logger.log_error(f"Failed to load FunASR model: {e}") + if self.on_error_callback: + await self.on_error_callback(str(e)) + raise + + async def disconnect(self) -> None: + """Clean up resources""" + self.is_connected_flag = False + self.model = None + self.audio_buffer.clear() + if self.logger: + self.logger.log_info("FunASR client disconnected") + + def is_connected(self) -> bool: + """Check if model is loaded""" + return self.is_connected_flag and self.model is not None + + async def send_audio(self, audio_data: bytes) -> None: + """Add audio data to buffer and process once enough has accumulated.""" + if not self.is_connected(): + return + + self.audio_buffer.extend(audio_data) + + # Calculate buffer duration in milliseconds (16-bit PCM => 2 bytes/sample) + buffer_duration_ms = ( + len(self.audio_buffer) / (self.sample_rate * 2) * 1000 + ) + + if buffer_duration_ms >= self.min_audio_length_ms: + await self._process_audio() + + async def _process_audio(self) -> None: + """Process accumulated audio buffer with the FunASR model.""" + async with self.processing_lock: + if len(self.audio_buffer) == 0: + return + + try: + # 16-bit PCM bytes -> float32 in [-1, 1] + audio_np = ( + np.frombuffer(self.audio_buffer, dtype=np.int16).astype( + np.float32 + ) + / 32768.0 + ) + + # Limit to max length + max_samples = int( + self.max_audio_length_ms / 1000 * self.sample_rate + ) + if len(audio_np) > max_samples: + audio_np = audio_np[:max_samples] + + duration_ms = int(len(audio_np) / self.sample_rate * 1000) + + # Run inference in a thread pool to avoid blocking the loop + loop = asyncio.get_event_loop() + res = await loop.run_in_executor( + None, + lambda: self.model.generate( + input=audio_np, + language=self.language, + use_itn=self.use_itn, + ), + ) + + # SenseVoice output carries tags like <|zh|><|NEUTRAL|>...; strip them. + text = ( + rich_transcription_postprocess(res[0]["text"]).strip() + if res + else "" + ) + + if text and self.on_result_callback: + await self.on_result_callback( + text=text, + start_ms=0, + duration_ms=duration_ms, + language=self.language, + final=True, + ) + + # Clear processed audio + self.audio_buffer.clear() + + except Exception as e: + # Clear buffer on error to prevent accumulation + self.audio_buffer.clear() + if self.logger: + self.logger.log_error(f"Error processing audio: {e}") + if self.on_error_callback: + await self.on_error_callback(str(e)) + + async def finalize(self) -> None: + """Process any remaining audio in buffer""" + if len(self.audio_buffer) > 0: + await self._process_audio() diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json b/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json new file mode 100644 index 0000000000..a7a70b3238 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json @@ -0,0 +1,57 @@ +{ + "type": "extension", + "name": "funasr_asr_python", + "version": "0.1.0", + "dependencies": [ + { + "type": "system", + "name": "ten_runtime_python", + "version": "0.11" + }, + { + "type": "system", + "name": "ten_ai_base", + "version": "0.7" + } + ], + "api": { + "interface": [ + { + "import_uri": "../../system/ten_ai_base/api/asr-interface.json" + } + ], + "property": { + "properties": { + "params": { + "type": "object", + "properties": { + "model": { + "type": "string" + }, + "device": { + "type": "string" + }, + "language": { + "type": "string" + }, + "use_itn": { + "type": "bool" + }, + "sample_rate": { + "type": "int64" + } + } + } + } + } + }, + "package": { + "include": [ + "manifest.json", + "property.json", + "**.py", + "requirements.txt", + "README.md" + ] + } +} diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/property.json b/ai_agents/agents/ten_packages/extension/funasr_asr_python/property.json new file mode 100644 index 0000000000..29c90b0ea1 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/property.json @@ -0,0 +1,13 @@ +{ + "dump": false, + "dump_path": "/tmp", + "finalize_mode": "disconnect", + "silence_duration_ms": 1000, + "params": { + "model": "iic/SenseVoiceSmall", + "device": "cpu", + "language": "auto", + "use_itn": true, + "sample_rate": 16000 + } +} diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/reconnect_manager.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/reconnect_manager.py new file mode 100644 index 0000000000..f5760806c1 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/reconnect_manager.py @@ -0,0 +1,90 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +import asyncio +from typing import Callable, Optional +from ten_ai_base.message import ModuleError, ModuleErrorCode + + +class ReconnectManager: + """Manages reconnection attempts with exponential backoff""" + + def __init__( + self, + max_attempts: int = 5, + base_delay: float = 0.5, + logger: Optional[any] = None, + ): + self.max_attempts = max_attempts + self.base_delay = base_delay + self.current_attempts = 0 + self.logger = logger + + def can_retry(self) -> bool: + """Check if more retry attempts are available""" + return self.current_attempts < self.max_attempts + + def mark_connection_successful(self) -> None: + """Reset retry counter after successful connection""" + self.current_attempts = 0 + + def get_attempts_info(self) -> str: + """Get current attempts information""" + return f"{self.current_attempts}/{self.max_attempts}" + + async def handle_reconnect( + self, + connection_func: Callable, + error_handler: Optional[Callable] = None, + ) -> bool: + """Handle reconnection with exponential backoff + + Args: + connection_func: Async function to establish connection + error_handler: Optional async function to handle errors + + Returns: + bool: True if reconnection initiated, False if max attempts reached + """ + if not self.can_retry(): + if self.logger: + self.logger.log_error( + f"Max reconnection attempts ({self.max_attempts}) reached" + ) + if error_handler: + await error_handler( + ModuleError( + module="asr", + code=ModuleErrorCode.FATAL_ERROR.value, + message=f"Max reconnection attempts ({self.max_attempts}) reached", + ) + ) + return False + + self.current_attempts += 1 + delay = self.base_delay * (2 ** (self.current_attempts - 1)) + + if self.logger: + self.logger.log_info( + f"Reconnecting in {delay}s (attempt {self.current_attempts}/{self.max_attempts})" + ) + + await asyncio.sleep(delay) + + try: + await connection_func() + return True + except Exception as e: + if self.logger: + self.logger.log_error(f"Reconnection attempt failed: {e}") + if error_handler: + await error_handler( + ModuleError( + module="asr", + code=ModuleErrorCode.NON_FATAL_ERROR.value, + message=f"Reconnection attempt {self.current_attempts} failed: {str(e)}", + ) + ) + return False diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt b/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt new file mode 100644 index 0000000000..87b6137477 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt @@ -0,0 +1,4 @@ +funasr>=1.1.0 +numpy>=1.24.0 +pydantic>=2.0.0 +pytest==8.3.4 diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/__init__.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/__init__.py new file mode 100644 index 0000000000..b8c07eef1c --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/__init__.py @@ -0,0 +1,4 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_config.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_config.py new file mode 100644 index 0000000000..b634364c1a --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_config.py @@ -0,0 +1,45 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +import pytest +from funasr_asr_python.config import FunASRConfig + + +def test_config_default_values(): + """Test default configuration values""" + config = FunASRConfig() + assert config.dump is False + assert config.dump_path == "/tmp" + assert config.finalize_mode == "disconnect" + assert config.silence_duration_ms == 1000 + assert config.params == {} + + +def test_config_from_json(): + """Test configuration from JSON""" + json_str = """{ + "dump": true, + "dump_path": "/var/log", + "finalize_mode": "silence", + "params": { + "model": "iic/SenseVoiceSmall", + "device": "cpu", + "language": "auto" + } + }""" + config = FunASRConfig.model_validate_json(json_str) + assert config.dump is True + assert config.dump_path == "/var/log" + assert config.finalize_mode == "silence" + assert config.params["model"] == "iic/SenseVoiceSmall" + assert config.params["device"] == "cpu" + + +def test_normalized_language(): + """Language codes map to normalized BCP-47-ish tags.""" + config = FunASRConfig(params={"language": "zh"}) + assert config.normalized_language == "zh-CN" + config2 = FunASRConfig(params={"language": "auto"}) + assert config2.normalized_language == "auto" From 7a468a0795729bfc702c5d9c0b86a2353508f633 Mon Sep 17 00:00:00 2001 From: zhifu gao Date: Thu, 16 Jul 2026 23:13:50 +0000 Subject: [PATCH 2/2] fix: harden FunASR extension lifecycle and tests --- .../extension/funasr_asr_python/config.py | 17 ++- .../extension/funasr_asr_python/extension.py | 40 ++++-- .../funasr_asr_python/funasr_client.py | 57 +++++--- .../extension/funasr_asr_python/manifest.json | 1 + .../funasr_asr_python/pyproject.toml | 10 ++ .../funasr_asr_python/requirements.txt | 2 + .../funasr_asr_python/tests/bin/start | 9 ++ .../funasr_asr_python/tests/test_client.py | 85 ++++++++++++ .../funasr_asr_python/tests/test_extension.py | 131 ++++++++++++++++++ 9 files changed, 315 insertions(+), 37 deletions(-) create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/pyproject.toml create mode 100755 ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/bin/start create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_client.py create mode 100644 ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_extension.py diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py index 4b14bb0de0..977e48466f 100644 --- a/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/config.py @@ -40,9 +40,8 @@ def to_json(self, sensitive_handling: bool = False) -> str: ) return str(config_dict) - @property - def normalized_language(self) -> str: - """Convert language code to normalized format""" + def normalize_language(self, detected_language: str = "") -> str: + """Convert configured or detected language to a BCP-47-ish tag.""" language_map = { "zh": "zh-CN", "en": "en-US", @@ -59,5 +58,15 @@ def normalized_language(self) -> str: "ar": "ar-AE", } params_dict = self.params or {} - language_code = params_dict.get("language", "") or "" + configured_language = params_dict.get("language", "") or "" + language_code = ( + detected_language + if configured_language in {"", "auto"} and detected_language + else configured_language + ) return language_map.get(language_code, language_code) + + @property + def normalized_language(self) -> str: + """Convert the configured language code to a BCP-47-ish tag.""" + return self.normalize_language() diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py index 8d3d4a4594..0bb4d08585 100644 --- a/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/extension.py @@ -128,6 +128,8 @@ async def start_connection(self) -> None: if self.reconnect_manager: self.reconnect_manager.mark_connection_successful() + await self.on_connected() + # Reset timeline self.sent_user_audio_duration_ms_before_last_reset += ( self.audio_timeline.get_total_user_audio_duration() @@ -141,6 +143,7 @@ async def start_connection(self) -> None: except Exception as e: self.ten_env.log_error(f"Failed to start FunASR connection: {e}") + self.client = None await self.send_asr_error( ModuleError( module=MODULE_NAME_ASR, @@ -148,6 +151,15 @@ async def start_connection(self) -> None: message=str(e), ), ) + await self.on_disconnected( + code=ModuleErrorCode.FATAL_ERROR.value, + message=str(e), + vendor_info=ModuleErrorVendorInfo( + vendor=self.vendor(), + code="model_load_failed", + message=str(e), + ), + ) @override async def stop_connection(self) -> None: @@ -157,6 +169,7 @@ async def stop_connection(self) -> None: if self.client: await self.client.disconnect() self.client = None + await self.on_disconnected(code=0, message="stopped") self.ten_env.log_info("FunASR connection stopped") except Exception as e: self.ten_env.log_error(f"Error stopping FunASR connection: {e}") @@ -185,6 +198,7 @@ async def send_audio( if not self.is_connected() or not self.client: return False + buf = None try: buf = frame.lock_buf() audio_data = bytes(buf) @@ -196,13 +210,14 @@ async def send_audio( # Send to client await self.client.send_audio(audio_data) - frame.unlock_buf(buf) return True except Exception as e: self.ten_env.log_error(f"Error sending audio to FunASR: {e}") - frame.unlock_buf(buf) return False + finally: + if buf is not None: + frame.unlock_buf(buf) @override async def finalize(self, _session_id: str | None) -> None: @@ -214,13 +229,16 @@ async def finalize(self, _session_id: str | None) -> None: f"FunASR finalize start at {self.last_finalize_timestamp}" ) - finalize_mode = self.config.finalize_mode - if finalize_mode == "disconnect": - await self._handle_finalize_disconnect() - elif finalize_mode == "silence": - await self._handle_finalize_silence() - else: - raise ValueError(f"invalid finalize mode: {finalize_mode}") + try: + finalize_mode = self.config.finalize_mode + if finalize_mode == "disconnect": + await self._handle_finalize_disconnect() + elif finalize_mode == "silence": + await self._handle_finalize_silence() + else: + raise ValueError(f"invalid finalize mode: {finalize_mode}") + finally: + await self._finalize_end() async def _handle_finalize_disconnect(self) -> None: """Handle disconnect mode finalization""" @@ -295,7 +313,9 @@ async def _on_result( start_ms=actual_start_ms, duration_ms=duration_ms, language=( - self.config.normalized_language if self.config else language + self.config.normalize_language(language) + if self.config + else language ), words=[], ) diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py index e9bd769a9e..86a737eb01 100644 --- a/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/funasr_client.py @@ -4,8 +4,10 @@ # import asyncio -import numpy as np +import re from typing import Optional, Callable + +import numpy as np from funasr import AutoModel from funasr.utils.postprocess_utils import rich_transcription_postprocess @@ -35,6 +37,7 @@ def __init__( self.model: Optional[AutoModel] = None self.audio_buffer = bytearray() + self.processed_audio_duration_ms = 0 self.is_connected_flag = False self.processing_lock = asyncio.Lock() @@ -68,8 +71,6 @@ async def connect(self) -> None: self.is_connected_flag = False if self.logger: self.logger.log_error(f"Failed to load FunASR model: {e}") - if self.on_error_callback: - await self.on_error_callback(str(e)) raise async def disconnect(self) -> None: @@ -77,6 +78,7 @@ async def disconnect(self) -> None: self.is_connected_flag = False self.model = None self.audio_buffer.clear() + self.processed_audio_duration_ms = 0 if self.logger: self.logger.log_info("FunASR client disconnected") @@ -106,22 +108,28 @@ async def _process_audio(self) -> None: return try: + max_samples = max( + 1, + int(self.max_audio_length_ms / 1000 * self.sample_rate), + ) + sample_count = min(len(self.audio_buffer) // 2, max_samples) + if sample_count == 0: + return + byte_count = sample_count * 2 + audio_bytes = bytes(self.audio_buffer[:byte_count]) + del self.audio_buffer[:byte_count] + # 16-bit PCM bytes -> float32 in [-1, 1] audio_np = ( - np.frombuffer(self.audio_buffer, dtype=np.int16).astype( + np.frombuffer(audio_bytes, dtype=np.int16).astype( np.float32 ) / 32768.0 ) - # Limit to max length - max_samples = int( - self.max_audio_length_ms / 1000 * self.sample_rate - ) - if len(audio_np) > max_samples: - audio_np = audio_np[:max_samples] - duration_ms = int(len(audio_np) / self.sample_rate * 1000) + start_ms = self.processed_audio_duration_ms + self.processed_audio_duration_ms += duration_ms # Run inference in a thread pool to avoid blocking the loop loop = asyncio.get_event_loop() @@ -135,32 +143,35 @@ async def _process_audio(self) -> None: ) # SenseVoice output carries tags like <|zh|><|NEUTRAL|>...; strip them. - text = ( - rich_transcription_postprocess(res[0]["text"]).strip() - if res - else "" - ) + result = res[0] if res else {} + raw_text = result.get("text", "") + text = rich_transcription_postprocess(raw_text).strip() + detected_language = self._extract_language(raw_text, result) if text and self.on_result_callback: await self.on_result_callback( text=text, - start_ms=0, + start_ms=start_ms, duration_ms=duration_ms, - language=self.language, + language=detected_language or self.language, final=True, ) - # Clear processed audio - self.audio_buffer.clear() - except Exception as e: - # Clear buffer on error to prevent accumulation - self.audio_buffer.clear() if self.logger: self.logger.log_error(f"Error processing audio: {e}") if self.on_error_callback: await self.on_error_callback(str(e)) + @staticmethod + def _extract_language(raw_text: str, result: dict) -> str: + language = result.get("language") or result.get("lang") + if language: + return str(language) + + match = re.match(r"^<\|([a-z]{2,3})\|>", raw_text) + return match.group(1) if match else "" + async def finalize(self) -> None: """Process any remaining audio in buffer""" if len(self.audio_buffer) > 0: diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json b/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json index a7a70b3238..f2f1ab3ffe 100644 --- a/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/manifest.json @@ -50,6 +50,7 @@ "manifest.json", "property.json", "**.py", + "pyproject.toml", "requirements.txt", "README.md" ] diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/pyproject.toml b/ai_agents/agents/ten_packages/extension/funasr_asr_python/pyproject.toml new file mode 100644 index 0000000000..bcf403b8a8 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/pyproject.toml @@ -0,0 +1,10 @@ +[project] +name = "funasr-asr-python" +version = "0.1.0" +requires-python = ">=3.10" +dependencies = [ + "funasr>=1.1.0", + "numpy>=1.24.0", + "pydantic>=2.0.0", + "typing-extensions>=4.5.0", +] diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt b/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt index 87b6137477..0d188eb61f 100644 --- a/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/requirements.txt @@ -2,3 +2,5 @@ funasr>=1.1.0 numpy>=1.24.0 pydantic>=2.0.0 pytest==8.3.4 +pytest-asyncio>=0.23.0 +typing-extensions>=4.5.0 diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/bin/start b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/bin/start new file mode 100755 index 0000000000..8e78210572 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/bin/start @@ -0,0 +1,9 @@ +#!/bin/bash + +set -e + +cd "$(dirname "${BASH_SOURCE[0]}")/../.." + +export PYTHONPATH=.ten/app:.ten/app/ten_packages/system/ten_runtime_python/lib:.ten/app/ten_packages/system/ten_runtime_python/interface:.ten/app/ten_packages/system/ten_ai_base/interface:$PYTHONPATH + +pytest -s tests/ "$@" diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_client.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_client.py new file mode 100644 index 0000000000..acb817db2d --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_client.py @@ -0,0 +1,85 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +from unittest.mock import AsyncMock, MagicMock, patch + +import numpy as np +import pytest + +from funasr_asr_python.funasr_client import FunASRClient + + +def pcm_samples(count: int) -> bytes: + return np.zeros(count, dtype=np.int16).tobytes() + + +@pytest.fixture +def client() -> FunASRClient: + result_callback = AsyncMock() + instance = FunASRClient( + sample_rate=10, + on_result_callback=result_callback, + ) + instance.model = MagicMock() + instance.model.generate.return_value = [{"text": "hello"}] + instance.is_connected_flag = True + return instance + + +@pytest.mark.asyncio +async def test_process_audio_reports_cumulative_start_offset( + client: FunASRClient, +) -> None: + client.audio_buffer.extend(pcm_samples(10)) + await client._process_audio() + client.audio_buffer.extend(pcm_samples(10)) + await client._process_audio() + + assert client.on_result_callback.await_args_list[0].kwargs["start_ms"] == 0 + assert ( + client.on_result_callback.await_args_list[1].kwargs["start_ms"] == 1000 + ) + + +@pytest.mark.asyncio +async def test_process_audio_preserves_samples_beyond_max_chunk( + client: FunASRClient, +) -> None: + client.max_audio_length_ms = 1000 + client.audio_buffer.extend(pcm_samples(15)) + + await client._process_audio() + + assert bytes(client.audio_buffer) == pcm_samples(5) + + +@pytest.mark.asyncio +async def test_process_audio_extracts_sensevoice_language( + client: FunASRClient, +) -> None: + client.model.generate.return_value = [ + {"text": "<|zh|><|NEUTRAL|><|Speech|><|woitn|>你好"} + ] + client.audio_buffer.extend(pcm_samples(10)) + + await client._process_audio() + + assert client.on_result_callback.await_args.kwargs["language"] == "zh" + + +@pytest.mark.asyncio +async def test_connect_failure_is_reported_by_extension_only() -> None: + error_callback = AsyncMock() + client = FunASRClient(on_error_callback=error_callback) + + with patch( + "funasr_asr_python.funasr_client.AutoModel", + side_effect=RuntimeError("load failed"), + ): + with pytest.raises(RuntimeError, match="load failed"): + await client.connect() + + error_callback.assert_not_awaited() + assert not client.is_connected() diff --git a/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_extension.py b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_extension.py new file mode 100644 index 0000000000..1f2c8652f2 --- /dev/null +++ b/ai_agents/agents/ten_packages/extension/funasr_asr_python/tests/test_extension.py @@ -0,0 +1,131 @@ +# +# This file is part of TEN Framework, an open source project. +# Licensed under the Apache License, Version 2.0. +# + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from funasr_asr_python.config import FunASRConfig +from funasr_asr_python.extension import FunASRExtension + + +@pytest.fixture +def extension() -> FunASRExtension: + instance = FunASRExtension("test_funasr") + instance.ten_env = MagicMock() + instance.ten_env.log_debug = MagicMock() + instance.ten_env.log_error = MagicMock() + instance.ten_env.log_info = MagicMock() + instance.ten_env.log_warn = MagicMock() + instance.ten_env.send_data = AsyncMock() + return instance + + +@pytest.mark.asyncio +async def test_start_connection_reports_connected( + extension: FunASRExtension, +) -> None: + extension.config = FunASRConfig(params={}) + extension.on_connected = AsyncMock() + extension.audio_timeline = MagicMock() + extension.audio_timeline.get_total_user_audio_duration.return_value = 0 + + with patch("funasr_asr_python.extension.FunASRClient") as client_class: + client = client_class.return_value + client.connect = AsyncMock() + client.is_connected.return_value = False + + await extension.start_connection() + + extension.on_connected.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_start_connection_failure_reports_disconnected( + extension: FunASRExtension, +) -> None: + extension.config = FunASRConfig(params={}) + extension.send_asr_error = AsyncMock() + extension.on_disconnected = AsyncMock() + + with patch("funasr_asr_python.extension.FunASRClient") as client_class: + client = client_class.return_value + client.connect = AsyncMock(side_effect=RuntimeError("load failed")) + client.is_connected.return_value = False + + await extension.start_connection() + + assert extension.send_asr_error.await_count == 1 + extension.on_disconnected.assert_awaited_once() + assert extension.client is None + + +@pytest.mark.asyncio +async def test_stop_connection_reports_disconnected( + extension: FunASRExtension, +) -> None: + client = MagicMock() + client.disconnect = AsyncMock() + extension.client = client + extension.on_disconnected = AsyncMock() + + await extension.stop_connection() + + extension.on_disconnected.assert_awaited_once_with( + code=0, message="stopped" + ) + + +@pytest.mark.asyncio +async def test_finalize_always_reports_completion( + extension: FunASRExtension, +) -> None: + extension.config = FunASRConfig(finalize_mode="disconnect") + client = MagicMock() + client.finalize = AsyncMock() + extension.client = client + extension.send_asr_finalize_end = AsyncMock() + + await extension.finalize(None) + + client.finalize.assert_awaited_once_with() + extension.send_asr_finalize_end.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_send_audio_does_not_mask_lock_failure( + extension: FunASRExtension, +) -> None: + client = MagicMock() + client.is_connected.return_value = True + extension.client = client + frame = MagicMock() + frame.lock_buf.side_effect = RuntimeError("lock failed") + + result = await extension.send_audio(frame, None) + + assert result is False + frame.unlock_buf.assert_not_called() + + +@pytest.mark.asyncio +async def test_result_uses_detected_language_when_config_is_auto( + extension: FunASRExtension, +) -> None: + extension.config = FunASRConfig(params={"language": "auto"}) + extension.audio_timeline = MagicMock() + extension.audio_timeline.get_audio_duration_before_time.return_value = 0 + extension.send_asr_result = AsyncMock() + + await extension._on_result( + text="你好", + start_ms=0, + duration_ms=1000, + language="zh", + final=False, + ) + + result = extension.send_asr_result.await_args.args[0] + assert result.language == "zh-CN"