server.py 8.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235
  1. import asyncio
  2. import base64
  3. import json
  4. import argparse
  5. import signal
  6. import sys
  7. import os
  8. from typing import Set
  9. import numpy as np
  10. import websockets
  11. from websockets.server import WebSocketServerProtocol
  12. # Force unbuffered output
  13. sys.stdout = os.fdopen(sys.stdout.fileno(), 'w', buffering=1)
  14. from .audio_processor import AudioBuffer
  15. try:
  16. from .qwen_engine import QwenASREngine
  17. QWEN_AVAILABLE = True
  18. except ImportError:
  19. QWEN_AVAILABLE = False
  20. try:
  21. from .engine import WhisperEngine
  22. WHISPER_AVAILABLE = True
  23. except ImportError:
  24. WHISPER_AVAILABLE = False
  25. class ASRServer:
  26. def __init__(
  27. self,
  28. model_type: str = "qwen",
  29. model_size: str = "small",
  30. port: int = 8765,
  31. host: str = "localhost",
  32. language: str = None,
  33. buffer_duration: float = 5.0
  34. ):
  35. self.port = port
  36. self.host = host
  37. self.buffer_duration = buffer_duration
  38. self.model_type = model_type
  39. self.clients: Set[WebSocketServerProtocol] = set()
  40. self.client_buffers: dict = {}
  41. self.is_running = True
  42. print(f"Initializing {model_type.upper()} ASR engine...")
  43. if model_type == "qwen" and QWEN_AVAILABLE:
  44. self.engine = QwenASREngine(
  45. # model_id="Qwen/Qwen3-ASR-1.7B",
  46. model_id="Qwen/Qwen3-ASR-0.6B",
  47. language=language
  48. )
  49. elif model_type == "whisper" and WHISPER_AVAILABLE:
  50. self.engine = WhisperEngine(
  51. model_size=model_size,
  52. language=language or "zh"
  53. )
  54. else:
  55. raise RuntimeError(f"Model type '{model_type}' not available. "
  56. f"Qwen: {QWEN_AVAILABLE}, Whisper: {WHISPER_AVAILABLE}")
  57. print("Engine loaded successfully")
  58. async def register(self, websocket: WebSocketServerProtocol):
  59. self.clients.add(websocket)
  60. self.client_buffers[id(websocket)] = AudioBuffer(
  61. max_duration=self.buffer_duration,
  62. sample_rate=16000
  63. )
  64. print(f"Client connected: {websocket.remote_address}")
  65. await self.send_message(websocket, {
  66. "type": "connected",
  67. "model": self.engine.get_model_info()
  68. })
  69. async def unregister(self, websocket: WebSocketServerProtocol):
  70. self.clients.remove(websocket)
  71. self.client_buffers.pop(id(websocket), None)
  72. print(f"Client disconnected: {websocket.remote_address}")
  73. async def send_message(self, websocket: WebSocketServerProtocol, message: dict):
  74. if websocket in self.clients:
  75. try:
  76. await websocket.send(json.dumps(message, ensure_ascii=False))
  77. except websockets.exceptions.ConnectionClosed:
  78. await self.unregister(websocket)
  79. async def handle_audio(self, websocket: WebSocketServerProtocol, data: bytes):
  80. ws_id = id(websocket)
  81. print(f"DEBUG: websocket id={ws_id}, type={type(ws_id)}")
  82. print(f"DEBUG: buffers keys={list(self.client_buffers.keys())}, types={[type(k) for k in self.client_buffers.keys()]}")
  83. client_buffer = self.client_buffers.get(ws_id)
  84. if not client_buffer:
  85. # Try direct access
  86. try:
  87. client_buffer = self.client_buffers[ws_id]
  88. print(f"Direct access worked!")
  89. except KeyError:
  90. print(f"ERROR: No buffer for client, websocket id={ws_id}")
  91. return
  92. print(f"Received audio data: {len(data)} bytes, buffer samples: {len(client_buffer.buffer)}")
  93. int16_array = np.frombuffer(data, dtype=np.int16)
  94. float_array = int16_array.astype(np.float32) / 32768.0
  95. client_buffer.append(float_array)
  96. print(f"Buffer now has {len(client_buffer.buffer)} samples")
  97. async def handle_message(self, websocket: WebSocketServerProtocol, message: dict):
  98. msg_type = message.get("type")
  99. print(f"Received message type: {msg_type}")
  100. if msg_type == "audio":
  101. audio_data = message.get("data")
  102. if audio_data:
  103. try:
  104. print(f"Decoding audio data, length: {len(audio_data)}")
  105. audio_bytes = base64.b64decode(audio_data)
  106. print(f"Audio decoded, bytes: {len(audio_bytes)}")
  107. await self.handle_audio(websocket, audio_bytes)
  108. except Exception as e:
  109. print(f"Error processing audio: {e}")
  110. import traceback
  111. traceback.print_exc()
  112. elif msg_type == "transcribe":
  113. client_buffer = self.client_buffers.get(id(websocket))
  114. if client_buffer:
  115. audio = client_buffer.get_last(self.buffer_duration)
  116. print(f"Transcribe request: {len(audio)} samples in buffer")
  117. if len(audio) > 1600:
  118. try:
  119. print("Starting transcription...")
  120. result = self.engine.transcribe_array(audio)
  121. print(f"Transcription result: {result['text']}")
  122. await self.send_message(websocket, {
  123. "type": "transcript",
  124. "text": result["text"],
  125. "language": result["language"],
  126. "segments": result.get("segments", [])
  127. })
  128. except Exception as e:
  129. print(f"Transcription error: {e}")
  130. import traceback
  131. traceback.print_exc()
  132. await self.send_message(websocket, {
  133. "type": "error",
  134. "message": str(e)
  135. })
  136. elif msg_type == "clear":
  137. client_buffer = self.client_buffers.get(id(websocket))
  138. if client_buffer:
  139. client_buffer.clear()
  140. await self.send_message(websocket, {"type": "cleared"})
  141. async def handler(self, websocket: WebSocketServerProtocol):
  142. await self.register(websocket)
  143. try:
  144. async for raw_message in websocket:
  145. try:
  146. if isinstance(raw_message, str):
  147. message = json.loads(raw_message)
  148. await self.handle_message(websocket, message)
  149. elif isinstance(raw_message, bytes):
  150. await self.handle_audio(websocket, raw_message)
  151. except json.JSONDecodeError:
  152. print(f"Invalid JSON from {websocket.remote_address}")
  153. except Exception as e:
  154. print(f"Error handling message: {e}")
  155. await self.send_message(websocket, {
  156. "type": "error",
  157. "message": str(e)
  158. })
  159. except websockets.exceptions.ConnectionClosed:
  160. pass
  161. finally:
  162. await self.unregister(websocket)
  163. async def start(self):
  164. print(f"Starting ASR server on {self.host}:{self.port}")
  165. async with websockets.serve(self.handler, self.host, self.port):
  166. await asyncio.Future()
  167. def run(self):
  168. asyncio.run(self.start())
  169. def main():
  170. parser = argparse.ArgumentParser(description="ASR WebSocket Server")
  171. parser.add_argument("--model", "-m", default="qwen",
  172. choices=["qwen", "whisper"],
  173. help="ASR model type")
  174. parser.add_argument("--size", "-s", default="small",
  175. choices=["tiny", "base", "small", "medium", "large"],
  176. help="Whisper model size (for whisper mode)")
  177. parser.add_argument("--port", "-p", type=int, default=8765,
  178. help="Server port")
  179. parser.add_argument("--host", default="localhost",
  180. help="Server host")
  181. parser.add_argument("--language", "-l", default=None,
  182. help="Target language code (None for auto)")
  183. parser.add_argument("--buffer", "-b", type=float, default=5.0,
  184. help="Audio buffer duration in seconds")
  185. args = parser.parse_args()
  186. server = ASRServer(
  187. model_type=args.model,
  188. model_size=args.size,
  189. port=args.port,
  190. host=args.host,
  191. language=args.language,
  192. buffer_duration=args.buffer
  193. )
  194. def signal_handler(sig, frame):
  195. print("\nShutting down server...")
  196. server.is_running = False
  197. sys.exit(0)
  198. signal.signal(signal.SIGINT, signal_handler)
  199. signal.signal(signal.SIGTERM, signal_handler)
  200. server.run()
  201. if __name__ == "__main__":
  202. main()