update:前端增加可配置的本地唤醒词运行时

This commit is contained in:
3030332422
2026-04-14 16:54:37 +08:00
parent 705b5732a6
commit d5b883817c
18 changed files with 1037 additions and 150 deletions
@@ -0,0 +1,250 @@
# Test 页面语音唤醒
## 概述
Test 页面集成了基于 **Sherpa-ONNX** 的高精度语音唤醒功能,支持自定义唤醒词和实时检测。使用轻量级关键词检测模型,提供毫秒级响应速度。
## 唤醒词模型
### 模型下载(必需)
**重要说明**: 项目不包含模型文件,需要提前下载配置。
### 官方模型下载地址
- **官方模型列表**: <https://csukuangfj.github.io/sherpa/onnx/kws/pretrained_models/index.html>
- **推荐模型**: `sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01`
### 下载和配置步骤
#### 1. 下载模型包
```bash
# 方法1:直接下载(推荐)
cd main/xiaozhi-server/test
wget https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01.tar.bz2
# 解压
tar xvf sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01.tar.bz2
# 方法2:使用ModelScope
pip install modelscope
python -c "
from modelscope import snapshot_download
snapshot_download('pkufool/sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01', cache_dir='./models')
"
```
#### 2. 配置模型文件
模型包下载后包含以下文件:
```
sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01/
├── encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx # 速度优先
├── encoder-epoch-12-avg-2-chunk-16-left-64.onnx #
├── encoder-epoch-99-avg-1-chunk-16-left-64.int8.onnx # 速度优先
├── encoder-epoch-99-avg-1-chunk-16-left-64.onnx # 精度优先
├── decoder-epoch-12-avg-2-chunk-16-left-64.onnx #
├── decoder-epoch-99-avg-1-chunk-16-left-64.onnx # 精度优先
├── joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx # 速度优先
├── joiner-epoch-12-avg-2-chunk-16-left-64.onnx #
├── joiner-epoch-99-avg-1-chunk-16-left-64.int8.onnx # 速度优先
├── joiner-epoch-99-avg-1-chunk-16-left-64.onnx # 精度优先
├── tokens.txt # Token映射表(必需)
├── keywords_raw.txt # 模型包里可能附带(可选,runtime 不依赖)
├── keywords.txt # 现成的
├── test_wavs/ # 测试音频(可选)
├── configuration.json # 模型元信息(可选)
└── README.md # 说明文档(可选)
```
#### 3. 选择配置方案
**方案一:精度优先(推荐)**
```bash
cd sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01
# 复制精度优先的epoch-99 fp32三件套
cp encoder-epoch-99-avg-1-chunk-16-left-64.onnx ../models/encoder.onnx
cp decoder-epoch-99-avg-1-chunk-16-left-64.onnx ../models/decoder.onnx
cp joiner-epoch-99-avg-1-chunk-16-left-64.onnx ../models/joiner.onnx
# 复制配套文件
cp tokens.txt ../models/tokens.txt
# keywords_raw.txt 如果模型包里附带,可自行保留;test runtime 不依赖它
```
**方案二:速度优先**
```bash
cd sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01
# 复制速度优先的epoch-99 int8三件套
cp encoder-epoch-99-avg-1-chunk-16-left-64.int8.onnx ../models/encoder.onnx
cp decoder-epoch-99-avg-1-chunk-16-left-64.onnx ../models/decoder.onnx
cp joiner-epoch-99-avg-1-chunk-16-left-64.int8.onnx ../models/joiner.onnx
# 复制配套文件
cp tokens.txt ../models/tokens.txt
```
**注意事项**:
- **不要混用fp32与int8**:三个模型文件必须保持一致的精度
- **优先选择epoch-99**:比epoch-12训练更充分,精度更高
- **必需文件**`encoder.onnx` + `decoder.onnx` + `joiner.onnx` + `tokens.txt` + `keywords.txt`
### 最终模型文件结构
配置完成后,你的models目录应该包含:
```
models/
├── encoder.onnx # 编码器模型(重命名后)
├── decoder.onnx # 解码器模型(重命名后)
├── joiner.onnx # 连接器模型(重命名后)
├── tokens.txt # 拼音Token映射表(228行版本)
├── keywords.txt # 关键词配置文件(需创建)
└── keywords_raw.txt # 可选,runtime 不依赖
```
## 启动方式
`main/xiaozhi-server/test` 目录执行:
```bash
python start_test_runtime.py
```
启动后默认地址:
- 页面地址:`http://127.0.0.1:8006/test_page.html`
- 事件桥地址:`ws://127.0.0.1:8006/wakeword-ws`
- 健康检查:`http://127.0.0.1:8006/health`
停止方式:
- 在运行终端按 `Ctrl+C`
- 会同时停止静态页面服务、事件桥和唤醒词检测流程
## 运行依赖
### 模型文件
模型目录至少需要这些文件:
- `encoder.onnx`
- `decoder.onnx`
- `joiner.onnx`
- `tokens.txt`
- `keywords.txt`
### Python 依赖
唤醒词服务依赖以下 Python 包:
- `sherpa-onnx`
- `sounddevice`
- `pypinyin`
**pip 安装(推荐)**
```bash
# 在 conda 环境中
conda activate xiaozhi-esp32-server
cd main/xiaozhi-server/test/wakeword_runtime/
pip install -r requirements.txt
```
> 如使用独立 conda 环境,可同时安装:
> ```bash
> conda create -n wakeword python=3.10 -y
> conda activate wakeword
> cd main/xiaozhi-server/test/wakeword_runtime/
> pip install -r requirements.txt
> ```
## 配置文件说明
配置文件位于 [config.json](./config.json)。
当前主要配置项:
```json
{
"wakeword": {
"enabled": true
},
"model_dir": "models",
"audio": {
"input_device": null,
"sample_rate": 16000,
"channels": 1
},
"detector": {
"num_threads": 4,
"provider": "cpu",
"max_active_paths": 2,
"keywords_score": 1.8,
"keywords_threshold": 0.1,
"num_trailing_blanks": 1,
"cooldown_seconds": 1.5
},
"logging": {
"level": "INFO",
"dir": "logs",
"file": "wakeword-runtime.log"
}
}
```
各字段含义:
| 参数 | 说明 |
| --- | --- |
| `wakeword.enabled` | 是否启用本地唤醒词检测 |
| `model_dir` | 模型和词表所在目录 |
| `audio.input_device` | 麦克风输入设备,默认使用系统默认设备 |
| `audio.sample_rate` | 采样率,默认 `16000` |
| `audio.channels` | 声道数,默认 `1` |
| `detector.num_threads` | 检测器线程数 |
| `detector.provider` | 推理 provider,当前通常为 `cpu` |
| `detector.max_active_paths` | 搜索路径数 |
| `detector.keywords_score` | 关键词增强分数 |
| `detector.keywords_threshold` | 检测阈值 |
| `detector.num_trailing_blanks` | 尾随空白数量 |
| `detector.cooldown_seconds` | 连续触发冷却时间 |
| `logging.level` | 日志等级 |
| `logging.dir` | 日志目录 |
| `logging.file` | 日志文件名 |
## 推荐使用流程
### 首次使用
1. 准备 `models/` 目录下的模型文件和 `tokens.txt`
2. 确认 `models/keywords.txt` 存在
3.`test` 目录运行 `python start_test_runtime.py`
4. 浏览器打开 `http://127.0.0.1:8006/test_page.html`
5. 进入设置页检查“唤醒词”配置
### 修改唤醒词
1. 打开 test 页面设置
2. 切到“唤醒词”页签
3. 修改启用状态或唤醒词列表
4. 点击“应用唤醒词”
5. 根据提示决定是否立即重启
### 禁用唤醒词
1. 将“启用本地唤醒词”改成禁用
2. 点击“应用唤醒词”
3. 建议立即重启一次
禁用后:
- 页面与事件桥仍然可用
- 唤醒词检测不会继续运行
@@ -19,12 +19,9 @@ class WakewordEventBridge:
return self._running
def build_ready_message(self) -> str:
return json.dumps(
{
"type": "bridge_connected",
"payload": {"status": "ready"},
},
ensure_ascii=False,
return self.build_message(
"bridge_connected",
{"status": "ready"},
)
def publish_detected(self, wake_word: str) -> None:
@@ -40,13 +37,7 @@ class WakewordEventBridge:
if not self._running:
return
message = json.dumps(
{
"type": event_type,
"payload": payload or {},
},
ensure_ascii=False,
)
message = self.build_message(event_type, payload or {})
with self._clients_lock:
clients = list(self._clients)
@@ -69,6 +60,24 @@ class WakewordEventBridge:
self._clients.append(client_queue)
return client_queue
def build_message(
self,
event_type: str,
payload: dict[str, Any] | None = None,
request_id: str | None = None,
success: bool = True,
error: str | None = None,
) -> str:
message: dict[str, Any] = {
"type": event_type,
"requestId": request_id,
"success": success,
"payload": payload or {},
}
if error:
message["error"] = error
return json.dumps(message, ensure_ascii=False)
def remove_client(self, client_queue: queue.Queue[str]) -> None:
with self._clients_lock:
if client_queue in self._clients:
@@ -2,7 +2,6 @@
"wakeword": {
"enabled": false
},
"wake_word": "你好小智",
"model_dir": "models",
"audio": {
"input_device": null,
@@ -14,7 +13,7 @@
"provider": "cpu",
"max_active_paths": 2,
"keywords_score": 1.8,
"keywords_threshold": 0.05,
"keywords_threshold": 0.1,
"num_trailing_blanks": 1,
"cooldown_seconds": 1.5
},
@@ -22,8 +21,5 @@
"level": "INFO",
"dir": "logs",
"file": "wakeword-runtime.log"
},
"wake_words": [
"你好小智"
]
}
}
}
@@ -1,7 +1,6 @@
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@dataclass
@@ -43,11 +42,10 @@ class RuntimeConfig:
audio: AudioSettings
detector: DetectorSettings
logging: LoggingSettings
raw: dict[str, Any]
def validate(self) -> None:
if self.wakeword.enabled and not self.wake_words:
raise ValueError("wake_word or wake_words cannot be empty when wakeword is enabled")
raise ValueError("keywords.txt cannot be empty when wakeword is enabled")
if self.audio.sample_rate <= 0:
raise ValueError("audio.sample_rate must be greater than 0")
@@ -85,17 +83,9 @@ def load_config(runtime_root: Path) -> RuntimeConfig:
config_path = runtime_root / "config.json"
raw = json.loads(config_path.read_text(encoding="utf-8"))
wakeword_cfg = dict(raw.get("wakeword", {}))
raw_words = raw.get("wake_words")
wake_words: list[str] = []
if isinstance(raw_words, list):
wake_words = [str(item).strip() for item in raw_words if str(item).strip()]
if not wake_words:
wake_word = str(raw.get("wake_word", "")).strip()
if wake_word:
wake_words = [wake_word]
raw_model_dir = Path(str(raw.get("model_dir", "models")))
model_dir = raw_model_dir.resolve() if raw_model_dir.is_absolute() else (runtime_root / raw_model_dir).resolve()
wake_words = _load_wake_words_from_keywords_file(model_dir)
audio_cfg = dict(raw.get("audio", {}))
detector_cfg = dict(raw.get("detector", {}))
@@ -125,7 +115,24 @@ def load_config(runtime_root: Path) -> RuntimeConfig:
directory=str(logging_cfg.get("dir", "logs")),
file_name=str(logging_cfg.get("file", "wakeword-runtime.log")),
),
raw=raw,
)
config.validate()
return config
def _load_wake_words_from_keywords_file(model_dir: Path) -> list[str]:
keywords_file = model_dir / "keywords.txt"
if not keywords_file.exists():
return []
wake_words: list[str] = []
for line in keywords_file.read_text(encoding="utf-8").splitlines():
text = line.strip()
if not text or text.startswith("#") or "@" not in text:
continue
wake_word = text.split("@", 1)[1].strip()
if wake_word:
wake_words.append(wake_word)
return wake_words
@@ -2,8 +2,6 @@ import queue
import threading
import logging
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
import numpy as np
@@ -14,14 +12,6 @@ from .detector_assets import DetectorAssetsBuilder
logger = logging.getLogger(__name__)
@dataclass
class DetectorBootstrapReport:
ready: bool
model_root: Path | None
keywords_file: Path | None
wake_words: list[str]
class WakewordDetector:
def __init__(self, config: RuntimeConfig) -> None:
self.config = config
@@ -39,7 +29,7 @@ class WakewordDetector:
self.last_detection_time = 0.0
self.detection_cooldown = self.config.detector.cooldown_seconds
def initialize(self) -> DetectorBootstrapReport:
def initialize(self) -> None:
if not self.enabled:
raise RuntimeError("wakeword detector is disabled")
@@ -69,17 +59,9 @@ class WakewordDetector:
provider=detector_cfg.provider,
)
self.stream = self.keyword_spotter.create_stream()
report = DetectorBootstrapReport(
ready=True,
model_root=assets.model_root,
keywords_file=assets.keywords_file,
wake_words=self.config.wake_words,
)
logger.info("detector initialized")
logger.info("detector model root: %s", assets.model_root)
logger.info("detector keywords file: %s", assets.keywords_file)
return report
def on_detected(self, callback: Callable[[str, str], None]) -> None:
self.on_detected_callback = callback
@@ -123,6 +105,9 @@ class WakewordDetector:
logger.info("wakeword detector started")
def stop(self) -> None:
if not self.is_running_flag and self.audio_source is None and self._worker_thread is None:
return
self.is_running_flag = False
if self.audio_source is not None:
@@ -56,7 +56,7 @@ class DetectorAssetsBuilder:
wake_words = self.config.wake_words
if not wake_words:
raise ValueError("wake_word or wake_words cannot be empty")
raise ValueError("keywords.txt cannot be empty")
keywords_path = model_root / "keywords.txt"
lines: list[str] = []
@@ -60,6 +60,9 @@ class MicrophoneListener:
logger.info("microphone block size: %s", self._block_size)
def stop(self) -> None:
if not self._running and self._stream is None:
return
self._running = False
if self._stream is not None:
try:
@@ -31,6 +31,7 @@ class AudioPlugin(Plugin):
self.source.stop()
def shutdown(self) -> None:
self.stop()
if self.app is not None:
self.app.audio_source = None
self.source = None
self.app = None
@@ -41,7 +41,8 @@ class WakeWordPlugin(Plugin):
self.detector.stop()
def shutdown(self) -> None:
self.stop()
self.detector = None
self.app = None
def _on_detected(self, wake_word: str, full_text: str) -> None:
if self.app is None:
@@ -0,0 +1,3 @@
sherpa-onnx==1.12.29
sounddevice>=0.4.4
pypinyin==0.55.0
@@ -1,17 +1,27 @@
import logging
from typing import Protocol
from ..config import RuntimeConfig
from ..plugins import AudioPlugin, PluginManager, WakeWordPlugin
from .http_server import TestRuntimeHttpServer
logger = logging.getLogger(__name__)
class EventPublisher(Protocol):
def publish_service_ready(self) -> None:
...
def publish_service_stopping(self) -> None:
...
def publish_detected(self, wake_word: str) -> None:
...
class TestRuntimeApplication:
def __init__(self, config: RuntimeConfig, http_server: TestRuntimeHttpServer) -> None:
def __init__(self, config: RuntimeConfig, event_publisher: EventPublisher) -> None:
self.config = config
self.http_server = http_server
self.event_bridge = http_server.event_bridge
self.event_bridge = event_publisher
self.plugins = PluginManager()
self.audio_source = None
self._is_setup = False
@@ -1,10 +1,16 @@
import base64
import hashlib
import json
import queue
import socket
import threading
from http import HTTPStatus
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Callable
from ..bridge import WakewordEventBridge
from ..config.config_loader import load_config
class TestRuntimeHttpServer:
@@ -13,6 +19,8 @@ class TestRuntimeHttpServer:
self.host = host
self.port = port
self.event_bridge = WakewordEventBridge()
self._restart_handler: Callable[[], None] | None = None
self._restart_lock = threading.Lock()
self._server = self._build_server()
@property
@@ -20,8 +28,8 @@ class TestRuntimeHttpServer:
return f"http://127.0.0.1:{self.port}/test_page.html"
@property
def events_url(self) -> str:
return f"http://127.0.0.1:{self.port}/events"
def bridge_url(self) -> str:
return f"ws://127.0.0.1:{self.port}/wakeword-ws"
def serve_forever(self) -> None:
self._server.serve_forever()
@@ -31,9 +39,32 @@ class TestRuntimeHttpServer:
self.event_bridge.close()
self._server.server_close()
def set_restart_handler(self, handler: Callable[[], None]) -> None:
self._restart_handler = handler
def request_runtime_restart(self) -> None:
with self._restart_lock:
handler = self._restart_handler
if handler is None:
raise RuntimeError("restart handler is not configured")
threading.Thread(
target=self._run_restart_handler,
name="test-runtime-restart",
daemon=True,
).start()
def _run_restart_handler(self) -> None:
handler = self._restart_handler
if handler is None:
return
handler()
def _build_server(self) -> ThreadingHTTPServer:
test_root = self.test_root
event_bridge = self.event_bridge
schedule_restart = self.request_runtime_restart
class TestRuntimeHandler(SimpleHTTPRequestHandler):
protocol_version = "HTTP/1.1"
@@ -48,8 +79,8 @@ class TestRuntimeHttpServer:
pass
def do_GET(self) -> None:
if self.path == "/events":
self._handle_events(event_bridge)
if self.path == "/wakeword-ws":
self._handle_websocket(event_bridge)
return
if self.path == "/health":
@@ -67,36 +98,255 @@ class TestRuntimeHttpServer:
def log_message(self, format: str, *args) -> None:
return
def _handle_events(self, bridge: WakewordEventBridge) -> None:
def _handle_websocket(self, bridge: WakewordEventBridge) -> None:
if self.headers.get("Upgrade", "").lower() != "websocket":
self.send_error(HTTPStatus.BAD_REQUEST, "expected websocket upgrade")
return
websocket_key = self.headers.get("Sec-WebSocket-Key")
if not websocket_key:
self.send_error(HTTPStatus.BAD_REQUEST, "missing Sec-WebSocket-Key")
return
accept_source = websocket_key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
accept_value = base64.b64encode(
hashlib.sha1(accept_source.encode("utf-8")).digest()
).decode("ascii")
client_queue = bridge.add_client()
self.send_response(HTTPStatus.OK)
self.send_header("Content-Type", "text/event-stream; charset=utf-8")
self.send_header("Cache-Control", "no-cache")
self.send_header("Connection", "keep-alive")
self.send_header("X-Accel-Buffering", "no")
self.send_response(HTTPStatus.SWITCHING_PROTOCOLS)
self.send_header("Upgrade", "websocket")
self.send_header("Connection", "Upgrade")
self.send_header("Sec-WebSocket-Accept", accept_value)
self.end_headers()
try:
ready_message = bridge.build_ready_message()
self.wfile.write(f"data: {ready_message}\n\n".encode("utf-8"))
self.wfile.flush()
self.connection.settimeout(0.2)
self._send_websocket_text(bridge.build_ready_message())
self._send_websocket_text(self._build_wakeword_config_message(bridge))
while bridge.is_running:
inbound_message = self._receive_websocket_message()
if inbound_message is not None:
response_message = self._handle_bridge_request(bridge, inbound_message)
if response_message:
self._send_websocket_text(response_message)
try:
message = client_queue.get(timeout=15)
message = client_queue.get(timeout=0.2)
if message == "__bridge_closed__":
break
self.wfile.write(f"data: {message}\n\n".encode("utf-8"))
self._send_websocket_text(message)
except queue.Empty:
if not bridge.is_running:
break
self.wfile.write(b": keepalive\n\n")
self.wfile.flush()
continue
except socket.timeout:
pass
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
finally:
bridge.remove_client(client_queue)
def _build_wakeword_config_message(self, bridge: WakewordEventBridge) -> str:
try:
runtime_root = test_root / "wakeword_runtime"
config = load_config(runtime_root)
payload = {
"enabled": config.wakeword_enabled,
"wakeWords": config.wake_words,
}
return bridge.build_message("wakeword_config", payload)
except Exception as exc:
return bridge.build_message(
"wakeword_config",
{},
success=False,
error=f"读取唤醒词配置失败: {exc}",
)
def _handle_bridge_request(self, bridge: WakewordEventBridge, raw_message: str) -> str | None:
try:
message = json.loads(raw_message)
except json.JSONDecodeError:
return None
message_type = str(message.get("type", "")).strip()
request_id = message.get("requestId")
payload = message.get("payload") or {}
result_type = f"{message_type}_result" if message_type else "bridge_request_result"
if message_type == "set_wakeword_config":
try:
result_payload = self._save_wakeword_config(payload)
bridge.publish("wakeword_config", result_payload)
return bridge.build_message(
"set_wakeword_config_result",
result_payload,
request_id=request_id,
)
except Exception as exc:
return bridge.build_message(
"set_wakeword_config_result",
{},
request_id=request_id,
success=False,
error=f"保存唤醒词配置失败: {exc}",
)
if message_type == "restart_wakeword_service":
schedule_restart()
return bridge.build_message(
"restart_wakeword_service_result",
{"restarting": True},
request_id=request_id,
)
return bridge.build_message(
result_type,
{},
request_id=request_id,
success=False,
error=f"unsupported message type: {message_type}",
)
def _save_wakeword_config(self, payload: dict) -> dict:
runtime_root = test_root / "wakeword_runtime"
config_path = runtime_root / "config.json"
model_root = runtime_root / "models"
keywords_path = model_root / "keywords.txt"
enabled = bool(payload.get("enabled", True))
wake_words = payload.get("wakeWords") or []
normalized_wake_words = []
for item in wake_words:
if not isinstance(item, str):
continue
text = item.strip()
if text and text not in normalized_wake_words:
normalized_wake_words.append(text)
if enabled and not normalized_wake_words:
raise ValueError("wakeWords cannot be empty when wakeword is enabled")
raw_config = json.loads(config_path.read_text(encoding="utf-8"))
raw_config.setdefault("wakeword", {})["enabled"] = enabled
config_path.write_text(
json.dumps(raw_config, indent=2, ensure_ascii=False),
encoding="utf-8",
)
keywords_lines = [self._build_keyword_line(item) for item in normalized_wake_words]
keywords_path.write_text(
("\n".join(keywords_lines) + "\n") if keywords_lines else "",
encoding="utf-8",
)
return {
"enabled": enabled,
"wakeWords": normalized_wake_words,
}
def _build_keyword_line(self, keyword_text: str) -> str:
from pypinyin import Style, pinyin
initials = pinyin(keyword_text, style=Style.INITIALS, strict=False)
finals = pinyin(
keyword_text,
style=Style.FINALS_TONE,
strict=False,
neutral_tone_with_five=True,
)
tokens: list[str] = []
for initial_parts, final_parts in zip(initials, finals):
initial = initial_parts[0].strip()
final = final_parts[0].strip()
if initial:
tokens.append(initial)
if final:
tokens.append(final)
if not tokens:
raise ValueError(f"failed to generate pinyin tokens for wake word: {keyword_text}")
return f"{' '.join(tokens)} @{keyword_text}"
def _receive_websocket_message(self) -> str | None:
try:
header = self._read_exact(2)
except socket.timeout:
return None
if not header:
return None
first_byte, second_byte = header[0], header[1]
opcode = first_byte & 0x0F
masked = (second_byte & 0x80) != 0
payload_length = second_byte & 0x7F
if payload_length == 126:
payload_length = int.from_bytes(self._read_exact(2), "big")
elif payload_length == 127:
payload_length = int.from_bytes(self._read_exact(8), "big")
masking_key = self._read_exact(4) if masked else b""
payload = self._read_exact(payload_length) if payload_length else b""
if masked and payload:
payload = bytes(
byte ^ masking_key[index % 4]
for index, byte in enumerate(payload)
)
if opcode == 0x8:
raise ConnectionAbortedError("websocket closed by client")
if opcode == 0x9:
self._send_websocket_frame(0xA, payload)
return None
if opcode == 0xA:
return None
if opcode != 0x1:
return None
return payload.decode("utf-8")
def _read_exact(self, size: int) -> bytes:
if size <= 0:
return b""
chunks = bytearray()
while len(chunks) < size:
chunk = self.connection.recv(size - len(chunks))
if not chunk:
raise ConnectionResetError("websocket connection closed")
chunks.extend(chunk)
return bytes(chunks)
def _send_websocket_text(self, message: str) -> None:
self._send_websocket_frame(0x1, message.encode("utf-8"))
def _send_websocket_frame(self, opcode: int, payload: bytes) -> None:
header = bytearray()
header.append(0x80 | opcode)
payload_length = len(payload)
if payload_length < 126:
header.append(payload_length)
elif payload_length < 65536:
header.append(126)
header.extend(payload_length.to_bytes(2, "big"))
else:
header.append(127)
header.extend(payload_length.to_bytes(8, "big"))
self.wfile.write(bytes(header) + payload)
self.wfile.flush()
server = ThreadingHTTPServer((self.host, self.port), TestRuntimeHandler)
server.daemon_threads = True
return server