mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-21 22:53:56 +08:00
353 lines
14 KiB
Python
353 lines
14 KiB
Python
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:
|
|
def __init__(self, test_root: Path, host: str = "0.0.0.0", port: int = 8006) -> None:
|
|
self.test_root = test_root
|
|
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
|
|
def page_url(self) -> str:
|
|
return f"http://127.0.0.1:{self.port}/test_page.html"
|
|
|
|
@property
|
|
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()
|
|
|
|
def shutdown(self) -> None:
|
|
self._server.shutdown()
|
|
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"
|
|
|
|
def __init__(self, *args, **kwargs):
|
|
super().__init__(*args, directory=str(test_root), **kwargs)
|
|
|
|
def handle(self) -> None:
|
|
try:
|
|
super().handle()
|
|
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
|
|
pass
|
|
|
|
def do_GET(self) -> None:
|
|
if self.path == "/wakeword-ws":
|
|
self._handle_websocket(event_bridge)
|
|
return
|
|
|
|
if self.path == "/health":
|
|
body = json.dumps({"status": "ok"}).encode("utf-8")
|
|
self.send_response(HTTPStatus.OK)
|
|
self.send_header("Content-Type", "application/json; charset=utf-8")
|
|
self.send_header("Cache-Control", "no-cache")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
return
|
|
|
|
super().do_GET()
|
|
|
|
def log_message(self, format: str, *args) -> None:
|
|
return
|
|
|
|
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.SWITCHING_PROTOCOLS)
|
|
self.send_header("Upgrade", "websocket")
|
|
self.send_header("Connection", "Upgrade")
|
|
self.send_header("Sec-WebSocket-Accept", accept_value)
|
|
self.end_headers()
|
|
|
|
try:
|
|
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=0.2)
|
|
if message == "__bridge_closed__":
|
|
break
|
|
self._send_websocket_text(message)
|
|
except queue.Empty:
|
|
if not bridge.is_running:
|
|
break
|
|
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
|