first commit

This commit is contained in:
2026-08-15 14:43:56 +08:00
commit 63b8c28e1d
340 changed files with 59515 additions and 0 deletions
+154
View File
@@ -0,0 +1,154 @@
from __future__ import annotations
import base64
import hashlib
import hmac
import json
import secrets
import time
from datetime import datetime
from pathlib import Path
from typing import Any
SESSION_COOKIE_NAME = "live_admin_session"
def _utc_now_iso() -> str:
return datetime.utcnow().replace(microsecond=0).isoformat() + "Z"
def _is_local_ip(ip: str) -> bool:
ip = (ip or "").strip()
return ip in {"127.0.0.1", "::1", "::ffff:127.0.0.1", "localhost"}
class AdminAuthManager:
def __init__(self, data_dir: str | Path, logger, *, session_ttl_sec: int = 12 * 60 * 60):
self.data_dir = Path(data_dir)
self.data_dir.mkdir(parents=True, exist_ok=True)
self.path = self.data_dir / "admin_auth.json"
self.audit_path = self.data_dir / "admin_audit.log"
self.logger = logger
self.session_ttl_sec = int(session_ttl_sec)
self._sessions: dict[str, dict[str, Any]] = {}
self._data = self._default_data()
self._load()
def _default_data(self) -> dict[str, Any]:
return {
"version": 1,
"created_at": "",
"password_hash": "",
"password_salt": "",
"password_updated_at": "",
}
def _load(self) -> None:
if not self.path.exists():
return
try:
data = json.loads(self.path.read_text(encoding="utf-8"))
if isinstance(data, dict):
self._data.update(data)
except Exception as exc:
self.logger.warning(f"[后台认证] 读取认证文件失败: {exc}")
def _save(self) -> None:
tmp = self.path.with_suffix(".tmp")
tmp.write_text(json.dumps(self._data, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
tmp.replace(self.path)
def _hash_password(self, password: str, salt: bytes) -> str:
digest = hashlib.scrypt(
password.encode("utf-8"),
salt=salt,
n=2 ** 14,
r=8,
p=1,
dklen=64,
)
return base64.b64encode(digest).decode("ascii")
def is_bootstrapped(self) -> bool:
return bool(self._data.get("password_hash") and self._data.get("password_salt"))
def can_bootstrap_ip(self, ip: str) -> bool:
return _is_local_ip(ip)
def bootstrap(self, password: str) -> None:
password = str(password or "")
if len(password) < 8:
raise ValueError("后台密码至少 8 位")
salt = secrets.token_bytes(16)
now = _utc_now_iso()
self._data = {
"version": 1,
"created_at": self._data.get("created_at") or now,
"password_hash": self._hash_password(password, salt),
"password_salt": base64.b64encode(salt).decode("ascii"),
"password_updated_at": now,
}
self._sessions.clear()
self._save()
def verify_password(self, password: str) -> bool:
if not self.is_bootstrapped():
return False
try:
salt = base64.b64decode(self._data["password_salt"])
except Exception:
return False
current = self._hash_password(str(password or ""), salt)
return hmac.compare_digest(current, self._data.get("password_hash", ""))
def create_session(self, client_ip: str, user_agent: str = "") -> tuple[str, int]:
self.cleanup_sessions()
token = secrets.token_urlsafe(32)
now = int(time.time())
expires_at = now + self.session_ttl_sec
self._sessions[token] = {
"created_at": now,
"expires_at": expires_at,
"client_ip": client_ip or "",
"user_agent": (user_agent or "")[:240],
}
return token, expires_at
def cleanup_sessions(self) -> None:
now = int(time.time())
expired = [token for token, session in self._sessions.items() if int(session.get("expires_at", 0)) <= now]
for token in expired:
self._sessions.pop(token, None)
def get_session(self, token: str) -> dict[str, Any] | None:
self.cleanup_sessions()
if not token:
return None
session = self._sessions.get(token)
if not session:
return None
session["expires_at"] = int(time.time()) + self.session_ttl_sec
return session
def destroy_session(self, token: str) -> None:
if token:
self._sessions.pop(token, None)
def write_audit(self, *, action: str, target: str = "", client_ip: str = "", session_id: str = "", detail: str = "") -> None:
entry = {
"at": _utc_now_iso(),
"action": action,
"target": target,
"client_ip": client_ip,
"session_id": session_id[:12],
"detail": detail,
}
line = json.dumps(entry, ensure_ascii=False)
try:
with open(self.audit_path, "a", encoding="utf-8") as handle:
handle.write(line + "\n")
except Exception as exc:
self.logger.warning(f"[后台认证] 写审计日志失败: {exc}")
self.logger.info(f"[审计] action={action} target={target or '-'} ip={client_ip or '-'} detail={detail or '-'}")
+43
View File
@@ -0,0 +1,43 @@
from __future__ import annotations
import asyncio
import time
class AdminEventBus:
def __init__(self):
self._subscribers: set[asyncio.Queue] = set()
self._counter = 0
self._lock = asyncio.Lock()
async def subscribe(self) -> asyncio.Queue:
queue: asyncio.Queue = asyncio.Queue(maxsize=64)
async with self._lock:
self._subscribers.add(queue)
return queue
async def unsubscribe(self, queue: asyncio.Queue) -> None:
async with self._lock:
self._subscribers.discard(queue)
def subscriber_count(self) -> int:
return len(self._subscribers)
async def publish(self, event_type: str, payload):
self._counter += 1
event = {
"id": self._counter,
"type": event_type,
"ts": time.time(),
"payload": payload,
}
for queue in list(self._subscribers):
if queue.full():
try:
queue.get_nowait()
except asyncio.QueueEmpty:
pass
try:
queue.put_nowait(event)
except asyncio.QueueFull:
pass
+601
View File
@@ -0,0 +1,601 @@
from __future__ import annotations
import asyncio
import base64
import ctypes
import ctypes.wintypes
import hashlib
import http.cookiejar
import http.cookies
import io
import json
import logging
import os
import secrets
import time
import urllib.parse
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any, Callable
_PUBLIC_KEY_PEM = """-----BEGIN PUBLIC KEY-----
MIGfMA0GCSqGSIb3DQEBAQUAA4GNADCBiQKBgQDLgd2OAkcGVtoE3ThUREbio0Eg
Uc/prcajMKXvkCKFCWhJYJcLkcM2DKKcSeFpD/j6Boy538YXnR6VhcuUJOhH2x71
nzPjfdTcqMz7djHum0qSZA0AyCBDABUqCrfNgCiJ00Ra7GmRj+YCK1NJEuewlb40
JNrRuoEUXpabUzGB8QIDAQAB
-----END PUBLIC KEY-----"""
_USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
_QR_GENERATE_URL = "https://passport.bilibili.com/x/passport-login/web/qrcode/generate"
_QR_POLL_URL = "https://passport.bilibili.com/x/passport-login/web/qrcode/poll"
_QR_HEADERS = {
"Referer": "https://www.bilibili.com/",
"Origin": "https://www.bilibili.com",
}
class _DataBlob(ctypes.Structure):
_fields_ = [("cbData", ctypes.wintypes.DWORD), ("pbData", ctypes.POINTER(ctypes.c_byte))]
def _blob(data: bytes) -> tuple[_DataBlob, Any]:
buffer = ctypes.create_string_buffer(data)
return _DataBlob(len(data), ctypes.cast(buffer, ctypes.POINTER(ctypes.c_byte))), buffer
def _dpapi_encrypt(value: str) -> str:
if os.name != "nt":
raise RuntimeError("B站刷新令牌安全存储仅支持 Windows DPAPI")
source, source_buffer = _blob(value.encode("utf-8"))
entropy, entropy_buffer = _blob(b"Live-streaming:bilibili-refresh-token:v1")
output = _DataBlob()
ok = ctypes.windll.crypt32.CryptProtectData(
ctypes.byref(source), None, ctypes.byref(entropy), None, None, 0,
ctypes.byref(output),
)
_ = source_buffer, entropy_buffer
if not ok:
raise ctypes.WinError()
try:
encrypted = ctypes.string_at(output.pbData, output.cbData)
return base64.b64encode(encrypted).decode("ascii")
finally:
ctypes.windll.kernel32.LocalFree(output.pbData)
def _dpapi_decrypt(value: str) -> str:
if os.name != "nt":
raise RuntimeError("B站刷新令牌安全存储仅支持 Windows DPAPI")
source, source_buffer = _blob(base64.b64decode(value))
entropy, entropy_buffer = _blob(b"Live-streaming:bilibili-refresh-token:v1")
output = _DataBlob()
ok = ctypes.windll.crypt32.CryptUnprotectData(
ctypes.byref(source), None, ctypes.byref(entropy), None, None, 0,
ctypes.byref(output),
)
_ = source_buffer, entropy_buffer
if not ok:
raise ctypes.WinError()
try:
return ctypes.string_at(output.pbData, output.cbData).decode("utf-8")
finally:
ctypes.windll.kernel32.LocalFree(output.pbData)
class BilibiliCredentialStore:
def __init__(self, path: str | Path):
self.path = Path(path)
def save_refresh_token(self, refresh_token: str) -> None:
token = str(refresh_token or "").strip()
if not token:
raise ValueError("refresh_token 不能为空")
payload = {
"version": 1,
"provider": "windows_dpapi_current_user",
"refresh_token_protected": _dpapi_encrypt(token),
}
self.path.parent.mkdir(parents=True, exist_ok=True)
tmp = self.path.with_suffix(self.path.suffix + ".tmp")
with open(tmp, "w", encoding="utf-8", newline="\n") as handle:
json.dump(payload, handle, ensure_ascii=False, indent=2)
handle.write("\n")
tmp.replace(self.path)
def load_refresh_token(self) -> str:
if not self.path.exists():
return ""
with open(self.path, "r", encoding="utf-8") as handle:
payload = json.load(handle)
protected = str(payload.get("refresh_token_protected") or "")
return _dpapi_decrypt(protected) if protected else ""
def is_configured(self) -> bool:
if not self.path.exists():
return False
try:
with open(self.path, "r", encoding="utf-8") as handle:
payload = json.load(handle)
return bool(str(payload.get("refresh_token_protected") or ""))
except (OSError, ValueError, TypeError):
return False
def _read_der_length(data: bytes, offset: int) -> tuple[int, int]:
first = data[offset]
offset += 1
if first < 0x80:
return first, offset
count = first & 0x7F
return int.from_bytes(data[offset:offset + count], "big"), offset + count
def _read_der_tlv(data: bytes, offset: int, expected_tag: int | None = None) -> tuple[int, bytes, int]:
tag = data[offset]
if expected_tag is not None and tag != expected_tag:
raise ValueError(f"DER tag 不匹配: expected={expected_tag:#x}, actual={tag:#x}")
length, content_offset = _read_der_length(data, offset + 1)
end = content_offset + length
return tag, data[content_offset:end], end
def _public_numbers() -> tuple[int, int]:
body = "".join(line for line in _PUBLIC_KEY_PEM.splitlines() if not line.startswith("-----"))
der = base64.b64decode(body)
_, spki, _ = _read_der_tlv(der, 0, 0x30)
_, _, offset = _read_der_tlv(spki, 0, 0x30)
_, bit_string, _ = _read_der_tlv(spki, offset, 0x03)
_, rsa_key, _ = _read_der_tlv(bit_string[1:], 0, 0x30)
_, modulus_bytes, rsa_offset = _read_der_tlv(rsa_key, 0, 0x02)
_, exponent_bytes, _ = _read_der_tlv(rsa_key, rsa_offset, 0x02)
return int.from_bytes(modulus_bytes, "big"), int.from_bytes(exponent_bytes, "big")
def _mgf1(seed: bytes, length: int) -> bytes:
result = bytearray()
counter = 0
while len(result) < length:
result.extend(hashlib.sha256(seed + counter.to_bytes(4, "big")).digest())
counter += 1
return bytes(result[:length])
def _rsa_oaep_sha256_encrypt(message: bytes) -> str:
modulus, exponent = _public_numbers()
key_size = (modulus.bit_length() + 7) // 8
digest_size = hashlib.sha256().digest_size
if len(message) > key_size - 2 * digest_size - 2:
raise ValueError("待加密内容过长")
label_hash = hashlib.sha256(b"").digest()
padding = b"\x00" * (key_size - len(message) - 2 * digest_size - 2)
data_block = label_hash + padding + b"\x01" + message
seed = secrets.token_bytes(digest_size)
data_mask = _mgf1(seed, key_size - digest_size - 1)
masked_data = bytes(left ^ right for left, right in zip(data_block, data_mask))
seed_mask = _mgf1(masked_data, digest_size)
masked_seed = bytes(left ^ right for left, right in zip(seed, seed_mask))
encoded = b"\x00" + masked_seed + masked_data
encrypted = pow(int.from_bytes(encoded, "big"), exponent, modulus)
return encrypted.to_bytes(key_size, "big").hex()
def _parse_cookie(cookie_text: str) -> dict[str, str]:
parsed = http.cookies.SimpleCookie()
parsed.load(str(cookie_text or "").replace("; ", ";"))
return {name: morsel.value for name, morsel in parsed.items()}
def _cookie_header(values: dict[str, str]) -> str:
return "; ".join(f"{name}={value}" for name, value in values.items() if value)
def _request_json(url: str, *, cookie: str = "", data: dict[str, str] | None = None,
opener: urllib.request.OpenerDirector | None = None,
headers: dict[str, str] | None = None,
retries: int = 0) -> tuple[dict, Any]:
body = urllib.parse.urlencode(data).encode("utf-8") if data is not None else None
request = urllib.request.Request(url, data=body, method="POST" if body is not None else "GET")
request.add_header("User-Agent", _USER_AGENT)
for name, value in (headers or {}).items():
request.add_header(name, value)
if cookie:
request.add_header("Cookie", cookie)
if body is not None:
request.add_header("Content-Type", "application/x-www-form-urlencoded")
retry_count = max(0, int(retries))
for attempt in range(retry_count + 1):
try:
response = (opener or urllib.request.build_opener()).open(request, timeout=15)
return json.loads(response.read().decode("utf-8")), response
except urllib.error.HTTPError:
raise
except (urllib.error.URLError, ConnectionError, TimeoutError, OSError) as exc:
if attempt >= retry_count:
raise RuntimeError("连接B站登录服务失败,请稍后重试") from exc
time.sleep(0.4 * (attempt + 1))
raise AssertionError("unreachable")
def _seed_cookie_jar(jar: http.cookiejar.CookieJar, values: dict[str, str]) -> None:
for name, value in values.items():
jar.set_cookie(http.cookiejar.Cookie(
version=0, name=name, value=value, port=None, port_specified=False,
domain=".bilibili.com", domain_specified=True, domain_initial_dot=True,
path="/", path_specified=True, secure=False, expires=None, discard=True,
comment=None, comment_url=None, rest={}, rfc2109=False,
))
def _jar_values(jar: http.cookiejar.CookieJar) -> dict[str, str]:
return {cookie.name: cookie.value for cookie in jar}
def _login_url_cookie_values(url: str) -> dict[str, str]:
values: dict[str, str] = {}
query = urllib.parse.urlparse(str(url or "")).query
for item in query.split("&"):
raw_name, separator, raw_value = item.partition("=")
if not separator:
continue
name = urllib.parse.unquote_plus(raw_name)
if name in {"SESSDATA", "bili_jct", "DedeUserID", "DedeUserID__ckMd5", "sid", "buvid3"}:
values[name] = raw_value
return values
def _render_qr_png(content: str) -> bytes:
try:
import qrcode
from qrcode.constants import ERROR_CORRECT_M
except ImportError as exc:
raise RuntimeError("缺少 qrcode 依赖,请重新安装 requirements.txt") from exc
qr = qrcode.QRCode(
version=None,
error_correction=ERROR_CORRECT_M,
box_size=8,
border=3,
)
qr.add_data(content)
qr.make(fit=True)
image = qr.make_image(fill_color="black", back_color="white")
output = io.BytesIO()
image.save(output, format="PNG")
return output.getvalue()
class BilibiliQrLogin:
"""服务端持有二维码密钥和 CookieJar,前端只获取二维码图片与状态。"""
_STATUS_MESSAGES = {
"idle": "尚未开始扫码登录",
"awaiting_scan": "请使用哔哩哔哩客户端扫码",
"awaiting_confirm": "已扫码,请在手机上确认登录",
"completed": "登录成功,Cookie 与刷新凭据已更新",
"expired": "二维码已过期,请重新生成",
"failed": "扫码登录失败",
}
def __init__(
self,
*,
credential_store: BilibiliCredentialStore,
update_cookie: Callable[[dict[str, str]], None],
logger: logging.Logger,
on_logged_in: Callable[[], Any] | None = None,
ttl_seconds: int = 180,
):
self.credential_store = credential_store
self.update_cookie = update_cookie
self.logger = logger
self.on_logged_in = on_logged_in
self.ttl_seconds = max(60, int(ttl_seconds))
self._lock = asyncio.Lock()
self._session: dict[str, Any] | None = None
def _snapshot(self) -> dict[str, Any]:
session = self._session or {}
state = str(session.get("state") or "idle")
expires_at = float(session.get("expires_at") or 0)
expires_in = max(0, int(expires_at - time.time())) if expires_at else 0
return {
"success": True,
"state": state,
"message": str(session.get("message") or self._STATUS_MESSAGES.get(state, "")),
"expires_in": expires_in,
"has_qr_image": state in {"awaiting_scan", "awaiting_confirm"} and expires_in > 0,
"credential_configured": self.credential_store.is_configured(),
"account": session.get("account"),
}
def snapshot(self) -> dict[str, Any]:
return self._snapshot()
def _start_sync(self) -> dict[str, Any]:
jar = http.cookiejar.CookieJar()
opener = urllib.request.build_opener(urllib.request.HTTPCookieProcessor(jar))
payload, _ = _request_json(
_QR_GENERATE_URL,
opener=opener,
headers=_QR_HEADERS,
retries=2,
)
if payload.get("code") != 0:
raise RuntimeError(f"B站二维码申请失败 code={payload.get('code')}")
data = payload.get("data") or {}
qr_url = str(data.get("url") or "").strip()
qr_key = str(data.get("qrcode_key") or "").strip()
if not qr_url or not qr_key:
raise RuntimeError("B站二维码响应缺少必要字段")
now = time.time()
return {
"state": "awaiting_scan",
"message": self._STATUS_MESSAGES["awaiting_scan"],
"created_at": now,
"expires_at": now + self.ttl_seconds,
"qr_url": qr_url,
"qr_key": qr_key,
"jar": jar,
"opener": opener,
"account": None,
}
async def start(self) -> dict[str, Any]:
async with self._lock:
try:
self._session = await asyncio.to_thread(self._start_sync)
except Exception as exc:
self._session = {
"state": "failed",
"message": str(exc) or type(exc).__name__,
"expires_at": 0,
}
self.logger.warning("[B站扫码登录] 二维码申请失败: %s", type(exc).__name__)
return self._snapshot()
self.logger.info("[B站扫码登录] 二维码已生成,等待扫码")
return self._snapshot()
def _poll_sync(self, session: dict[str, Any]) -> dict[str, Any]:
url = _QR_POLL_URL + "?" + urllib.parse.urlencode({"qrcode_key": session["qr_key"]})
payload, _ = _request_json(
url,
opener=session["opener"],
headers=_QR_HEADERS,
retries=2,
)
if payload.get("code") != 0:
return {"state": "failed", "message": f"B站扫码状态查询失败 code={payload.get('code')}"}
data = payload.get("data") or {}
status_code = int(data.get("code") or 0)
if status_code == 86101:
return {"state": "awaiting_scan", "message": self._STATUS_MESSAGES["awaiting_scan"]}
if status_code == 86090:
return {"state": "awaiting_confirm", "message": self._STATUS_MESSAGES["awaiting_confirm"]}
if status_code == 86038:
return {"state": "expired", "message": self._STATUS_MESSAGES["expired"]}
if status_code != 0:
return {"state": "failed", "message": str(data.get("message") or f"扫码失败 code={status_code}")}
refresh_token = str(data.get("refresh_token") or "").strip()
if not refresh_token:
return {"state": "failed", "message": "扫码成功响应缺少 refresh_token"}
cookie_values = _jar_values(session["jar"])
for name, value in _login_url_cookie_values(str(data.get("url") or "")).items():
cookie_values.setdefault(name, value)
if not cookie_values.get("SESSDATA") or not cookie_values.get("bili_jct"):
return {"state": "failed", "message": "扫码成功但登录 Cookie 不完整"}
cookie = _cookie_header(cookie_values)
nav, _ = _request_json(
"https://api.bilibili.com/x/web-interface/nav",
cookie=cookie,
headers={"Referer": "https://www.bilibili.com/"},
retries=2,
)
nav_data = nav.get("data") or {}
if nav.get("code") != 0 or not bool(nav_data.get("isLogin")):
return {"state": "failed", "message": "扫码 Cookie 登录验证失败"}
self.credential_store.save_refresh_token(refresh_token)
self.update_cookie(cookie_values)
return {
"state": "completed",
"message": self._STATUS_MESSAGES["completed"],
"account": {
"mid": str(nav_data.get("mid") or ""),
"uname": str(nav_data.get("uname") or "B站账号"),
},
}
async def poll(self) -> dict[str, Any]:
callback_needed = False
async with self._lock:
if not self._session:
return self._snapshot()
state = str(self._session.get("state") or "idle")
if state in {"completed", "expired", "failed"}:
return self._snapshot()
if time.time() >= float(self._session.get("expires_at") or 0):
self._session.update(state="expired", message=self._STATUS_MESSAGES["expired"])
return self._snapshot()
try:
result = await asyncio.to_thread(self._poll_sync, self._session)
except Exception as exc:
self.logger.warning("[B站扫码登录] 状态查询异常: %s", type(exc).__name__)
self._session.update(
state="failed",
message=str(exc) or f"扫码状态查询异常: {type(exc).__name__}",
)
return self._snapshot()
previous_state = state
self._session.update(result)
callback_needed = previous_state != "completed" and result.get("state") == "completed"
snapshot = self._snapshot()
if callback_needed:
self.logger.info("[B站扫码登录] 登录成功,Cookie 与刷新凭据已更新")
if self.on_logged_in:
callback_result = self.on_logged_in()
if asyncio.iscoroutine(callback_result):
await callback_result
return snapshot
async def qr_png(self) -> bytes:
async with self._lock:
if not self._session or self._session.get("state") not in {"awaiting_scan", "awaiting_confirm"}:
raise RuntimeError("当前没有可用的登录二维码")
if time.time() >= float(self._session.get("expires_at") or 0):
self._session.update(state="expired", message=self._STATUS_MESSAGES["expired"])
raise RuntimeError("登录二维码已过期")
content = str(self._session.get("qr_url") or "")
return await asyncio.to_thread(_render_qr_png, content)
class BilibiliCookieRefresher:
def __init__(
self,
*,
credential_store: BilibiliCredentialStore,
get_cookie: Callable[[], str],
update_cookie: Callable[[dict[str, str]], None],
logger: logging.Logger,
check_interval_seconds: int = 6 * 60 * 60,
on_refreshed: Callable[[], Any] | None = None,
is_enabled: Callable[[], bool] | None = None,
get_check_interval_seconds: Callable[[], int] | None = None,
):
self.credential_store = credential_store
self.get_cookie = get_cookie
self.update_cookie = update_cookie
self.logger = logger
self.check_interval_seconds = max(3600, int(check_interval_seconds))
self.on_refreshed = on_refreshed
self.is_enabled = is_enabled or (lambda: True)
self.get_check_interval_seconds = get_check_interval_seconds
self._stop = False
self._wake = asyncio.Event()
def stop(self) -> None:
self._stop = True
self._wake.set()
def wake(self) -> None:
self._wake.set()
def _current_interval(self) -> int:
if not self.get_check_interval_seconds:
return self.check_interval_seconds
try:
return max(3600, int(self.get_check_interval_seconds()))
except (TypeError, ValueError):
return self.check_interval_seconds
def _check_and_refresh_sync(self) -> dict[str, Any]:
refresh_token = self.credential_store.load_refresh_token()
if not refresh_token:
return {"status": "disabled", "message": "未配置刷新令牌"}
current_cookie = self.get_cookie()
current_values = _parse_cookie(current_cookie)
csrf = current_values.get("bili_jct", "")
if not current_values.get("SESSDATA") or not csrf:
return {"status": "failed", "message": "当前 Cookie 缺少 SESSDATA 或 bili_jct"}
info, _ = _request_json(
"https://passport.bilibili.com/x/passport-login/web/cookie/info?" +
urllib.parse.urlencode({"csrf": csrf}),
cookie=current_cookie,
retries=2,
)
if info.get("code") != 0:
return {"status": "failed", "message": f"登录状态检查失败 code={info.get('code')}"}
if not bool((info.get("data") or {}).get("refresh")):
return {"status": "valid", "message": "Cookie 当前无需刷新"}
timestamp = str((info.get("data") or {}).get("timestamp") or "")
correspond_path = _rsa_oaep_sha256_encrypt(f"refresh_{timestamp}".encode("utf-8"))
request = urllib.request.Request(
f"https://www.bilibili.com/correspond/1/{correspond_path}",
headers={"User-Agent": _USER_AGENT, "Cookie": current_cookie},
)
html = urllib.request.urlopen(request, timeout=15).read().decode("utf-8", errors="replace")
match = __import__("re").search(r'<div\s+id=["\']1-name["\']>([^<]+)</div>', html)
if not match:
return {"status": "failed", "message": "未获取到 refresh_csrf"}
refresh_csrf = match.group(1).strip()
old_refresh_token = refresh_token
jar = http.cookiejar.CookieJar()
_seed_cookie_jar(jar, current_values)
opener = urllib.request.build_opener(urllib.request.HTTPCookieProcessor(jar))
refreshed, _ = _request_json(
"https://passport.bilibili.com/x/passport-login/web/cookie/refresh",
data={
"csrf": csrf,
"refresh_csrf": refresh_csrf,
"source": "main_web",
"refresh_token": old_refresh_token,
},
opener=opener,
)
if refreshed.get("code") != 0:
return {"status": "failed", "message": f"Cookie 刷新失败 code={refreshed.get('code')}"}
new_refresh_token = str((refreshed.get("data") or {}).get("refresh_token") or "")
if not new_refresh_token:
return {"status": "failed", "message": "刷新响应缺少新 refresh_token"}
new_values = dict(current_values)
new_values.update(_jar_values(jar))
new_csrf = new_values.get("bili_jct", "")
new_cookie = _cookie_header(new_values)
confirmed, _ = _request_json(
"https://passport.bilibili.com/x/passport-login/web/confirm/refresh",
cookie=new_cookie,
data={"csrf": new_csrf, "refresh_token": old_refresh_token},
)
if confirmed.get("code") != 0:
return {"status": "failed", "message": f"刷新确认失败 code={confirmed.get('code')}"}
nav, _ = _request_json(
"https://api.bilibili.com/x/web-interface/nav",
cookie=new_cookie,
retries=2,
)
if nav.get("code") != 0 or not bool((nav.get("data") or {}).get("isLogin")):
return {"status": "failed", "message": "新 Cookie 登录验证失败"}
self.credential_store.save_refresh_token(new_refresh_token)
self.update_cookie(new_values)
return {"status": "refreshed", "message": "Cookie 已刷新并验证"}
async def check_once(self) -> dict[str, Any]:
try:
result = await asyncio.to_thread(self._check_and_refresh_sync)
except Exception as exc:
self.logger.warning("[B站凭据] 自动检查异常: %s", type(exc).__name__)
return {"status": "failed", "message": type(exc).__name__}
status = result.get("status")
if status == "refreshed":
self.logger.info("[B站凭据] Cookie 已自动续期并完成登录验证")
if self.on_refreshed:
callback_result = self.on_refreshed()
if asyncio.iscoroutine(callback_result):
await callback_result
elif status == "valid":
self.logger.info("[B站凭据] Cookie 有效,当前无需续期")
elif status == "disabled":
self.logger.warning("[B站凭据] 自动续期未启用:未配置刷新令牌")
else:
self.logger.warning("[B站凭据] 自动续期失败:%s", result.get("message", "未知错误"))
return result
async def run(self) -> None:
while not self._stop:
if self.is_enabled():
await self.check_once()
try:
await asyncio.wait_for(self._wake.wait(), timeout=self._current_interval())
self._wake.clear()
except asyncio.TimeoutError:
continue
+1
View File
@@ -0,0 +1 @@
"""Core helpers for the live streaming app."""
+39
View File
@@ -0,0 +1,39 @@
"""Runtime path helpers.
All mutable runtime data is resolved from the executable directory when the
program is frozen, and from the project root while running from source.
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
def app_root() -> Path:
if getattr(sys, "frozen", False):
return Path(sys.executable).resolve().parent
return Path(__file__).resolve().parents[2]
APP_ROOT = app_root()
CONFIG_DIR = APP_ROOT / "config"
DATA_DIR = APP_ROOT / "data"
WEB_DIR = APP_ROOT / "web"
LOG_DIR = APP_ROOT / "logs"
INTEGRATIONS_DIR = APP_ROOT / "integrations"
VENDOR_DIR = APP_ROOT / "vendor"
DOTS_TTS_SRC = VENDOR_DIR / "dots.tts-main" / "src"
def project_path(value: str | os.PathLike, *, base: Path | None = None) -> Path:
path = Path(value)
if path.is_absolute():
return path
return (base or APP_ROOT) / path
def ensure_runtime_dirs() -> None:
for path in (CONFIG_DIR, DATA_DIR, WEB_DIR, LOG_DIR):
path.mkdir(parents=True, exist_ok=True)
+654
View File
@@ -0,0 +1,654 @@
"""
BetterGI 弹幕联动主程序
========================
监听 B 站直播间弹幕 -> 关键词匹配 -> 调用 BetterGI.exe --startGroups 执行配置组
依赖: pip install websockets brotli
运行: python danmu_bettergi.py
"""
import asyncio
import json
import logging
import os
import struct
import subprocess
import sys
import time
import urllib.request
import zlib
from pathlib import Path
APP_DIR = Path(__file__).resolve().parent
if str(APP_DIR) not in sys.path:
sys.path.insert(0, str(APP_DIR))
from core.runtime_paths import APP_ROOT as PROJECT_ROOT, CONFIG_DIR, LOG_DIR, ensure_runtime_dirs, project_path
try:
import websockets
import brotli
except ImportError:
print("缺少依赖,请先运行: pip install websockets brotli")
sys.exit(1)
# ============== B站直播弹幕协议常量 ==============
HEADER_LEN = 16
OP_HEARTBEAT = 2 # 心跳请求
OP_HEARTBEAT_REPLY = 3 # 心跳响应(人气值)
OP_MESSAGE = 5 # 业务消息
OP_AUTH = 7 # 认证请求
OP_AUTH_REPLY = 8 # 认证响应
PROTO_JSON = 0 # 明文JSON
PROTO_ZLIB = 2 # zlib压缩
PROTO_BROTLI = 3 # brotli压缩
# ============== 弹幕服务器发现 ==============
def get_danmu_server(room_id: int) -> dict:
"""通过B站API获取弹幕服务器地址和token。优先使用host_server_list中的新服务器。"""
url = f"https://api.live.bilibili.com/room/v1/Danmu/getConf?room_id={room_id}"
req = urllib.request.Request(url)
req.add_header("User-Agent", "Mozilla/5.0")
resp = urllib.request.urlopen(req, timeout=10)
data = json.loads(resp.read())["data"]
token = data["token"]
host_list = data.get("host_server_list", [])
if host_list:
entry = host_list[0]
host = entry["host"]
port = entry.get("wss_port", 443)
else:
host = data["host"]
port = data.get("wss_port", 443)
return {
"ws_url": f"wss://{host}:{port}/sub",
"token": token,
"host": host,
"port": port,
}
def get_buvid3(sessdata: str = "") -> str:
"""通过B站finger/spi接口获取buvid3 (2024+协议认证必需)。"""
url = "https://api.bilibili.com/x/frontend/finger/spi"
req = urllib.request.Request(url)
req.add_header("User-Agent", "Mozilla/5.0")
if sessdata:
req.add_header("Cookie", f"SESSDATA={sessdata}")
try:
resp = urllib.request.urlopen(req, timeout=10)
data = json.loads(resp.read())
if data.get("code") == 0:
return data["data"]["b_3"]
except Exception:
pass
return ""
# ============== 配置管理 ==============
class Config:
def __init__(self, path: str):
self.path = Path(path)
self.data = {}
self.reload()
def reload(self):
with open(self.path, "r", encoding="utf-8") as f:
self.data = json.load(f)
@property
def room_id(self) -> int:
return int(self.data["bilibili"]["room_id"])
@property
def sessdata(self) -> str:
return self.data["bilibili"].get("sessdata", "")
@property
def bettergi_exe(self) -> str:
return self.data["bettergi"]["exe_path"]
@property
def bettergi_work_dir(self) -> str:
wd = self.data["bettergi"].get("work_dir", "")
return wd if wd else str(Path(self.bettergi_exe).parent)
@property
def default_cooldown(self) -> int:
return int(self.data["global"].get("default_cooldown", 30))
@property
def admin_uids(self) -> set:
return set(int(x) for x in self.data["global"].get("admin_uids", []))
@property
def rules(self) -> list:
return self.data.get("rules", [])
@property
def restart_mode(self) -> str:
"""重启模式: gentle(温和,跳过) / aggressive(激进,先杀再启)"""
return self.data["global"].get("restart_mode", "gentle")
# ============== 冷却管理 ==============
class CooldownManager:
"""按规则记录最后触发时间,防止同一指令被弹幕刷屏重复触发。"""
def __init__(self):
self._last_fire: dict = {} # keyword -> timestamp
def can_fire(self, keyword: str, cooldown: int) -> bool:
now = time.time()
last = self._last_fire.get(keyword, 0)
return (now - last) >= cooldown
def mark_fired(self, keyword: str):
self._last_fire[keyword] = time.time()
# ============== BetterGI 调用 ==============
class BetterGIRunner:
"""封装 BetterGI.exe --startGroups 调用。"""
def __init__(self, exe_path: str, work_dir: str, logger: logging.Logger,
restart_mode: str = "gentle"):
self.exe_path = exe_path
self.work_dir = work_dir
self.logger = logger
self.restart_mode = restart_mode # gentle / aggressive
self._busy = False # 标记是否正在执行任务
self._lock = asyncio.Lock()
# 缓存 cancelTaskHotkey (从 BetterGI Config.json 读)
self._cancel_hotkey = None
self._load_cancel_hotkey()
def _load_cancel_hotkey(self):
"""从 BetterGI 的 Config.json 读取取消任务热键。"""
try:
import os
cfg_path = os.path.join(self.work_dir, "User", "Config.json")
if os.path.exists(cfg_path):
with open(cfg_path, "r", encoding="utf-8") as f:
bgi_cfg = json.load(f)
hk = bgi_cfg.get("hotKeyConfig", {}).get("cancelTaskHotkey", "")
if hk:
self._cancel_hotkey = hk
self.logger.info(f"已读取 BetterGI 取消任务热键: {hk}")
except Exception as e:
self.logger.debug(f"读取 cancelTaskHotkey 失败(不影响使用): {e}")
async def _stop_bgi_gracefully(self):
"""优雅停止 BetterGI: 先按取消热键,再 taskkill 兜底。"""
# ① 模拟按取消热键 (让 BGI 内部收尾,保存进度)
if self._cancel_hotkey:
self.logger.info(f"[激进] 按取消热键: {self._cancel_hotkey}")
try:
await self._send_key(self._cancel_hotkey)
except Exception as e:
self.logger.warning(f"[激进] 按热键失败: {e}")
# 给 BetterGI 5 秒收尾时间
await asyncio.sleep(5)
# ② taskkill 强杀兜底
self.logger.info("[激进] taskkill /F /IM BetterGI.exe")
try:
proc = await asyncio.create_subprocess_exec(
"taskkill", "/F", "/IM", "BetterGI.exe",
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
# CREATE_NO_WINDOW = 0x08000000,避免弹黑窗
creationflags=0x08000000,
)
await proc.communicate()
except Exception as e:
self.logger.debug(f"[激进] taskkill 执行: {e}")
# 再等 1 秒让进程完全退出
await asyncio.sleep(1)
async def _send_key(self, key: str):
"""模拟按键。优先用 pywin32,其次 keyboard 库,都没有则跳过。"""
# 简单映射: BetterGI 的热键名通常是 "F9" "Ctrl+P" 这种
# 优先尝试 pywin32 (最稳定,不弹窗)
try:
import win32api
import win32con
# 简单支持单键 (F1-F24, 字母, 数字)
vk_map = {
"F1": win32con.VK_F1, "F2": win32con.VK_F2, "F3": win32con.VK_F3,
"F4": win32con.VK_F4, "F5": win32con.VK_F5, "F6": win32con.VK_F6,
"F7": win32con.VK_F7, "F8": win32con.VK_F8, "F9": win32con.VK_F9,
"F10": win32con.VK_F10, "F11": win32con.VK_F11, "F12": win32con.VK_F12,
"ESC": win32con.VK_ESCAPE, "ESCAPE": win32con.VK_ESCAPE,
}
# 处理 Ctrl+X Shift+X 这种组合键
parts = key.replace("+", " ").split()
main_key = parts[-1].upper()
ctrl = "CTRL" in [p.upper() for p in parts[:-1]]
shift = "SHIFT" in [p.upper() for p in parts[:-1]]
alt = "ALT" in [p.upper() for p in parts[:-1]]
vk = vk_map.get(main_key)
if vk is None and len(main_key) == 1:
vk = ord(main_key.upper()) # 字母键
if vk is None:
self.logger.warning(f"[激进] 不支持的热键: {key}, 跳过热键直接 taskkill")
return
if ctrl:
win32api.keybd_event(win32con.VK_CONTROL, 0, 0, 0)
if shift:
win32api.keybd_event(win32con.VK_SHIFT, 0, 0, 0)
if alt:
win32api.keybd_event(win32con.VK_MENU, 0, 0, 0)
win32api.keybd_event(vk, 0, 0, 0) # 按下
win32api.keybd_event(vk, 0, win32con.KEYEVENTF_KEYUP, 0) # 抬起
if ctrl:
win32api.keybd_event(win32con.VK_CONTROL, 0, win32con.KEYEVENTF_KEYUP, 0)
if shift:
win32api.keybd_event(win32con.VK_SHIFT, 0, win32con.KEYEVENTF_KEYUP, 0)
if alt:
win32api.keybd_event(win32con.VK_MENU, 0, win32con.KEYEVENTF_KEYUP, 0)
self.logger.info(f"[激进] 已模拟按键: {key}")
except ImportError:
self.logger.warning(
"[激进] 未安装 pywin32, 无法模拟热键。"
"可运行: pip install pywin32 (可选,不装则直接 taskkill)"
)
async def run_groups(self, groups: list) -> bool:
"""异步启动配置组。"""
async with self._lock:
if self._busy:
if self.restart_mode == "aggressive":
self.logger.warning(
f"[激进模式] BetterGI 正在执行任务,先杀再启: {groups}"
)
# 释放锁,执行停止 (停止可能耗时,不持有锁)
self._busy = False
else:
self.logger.warning(
f"[温和模式] BetterGI 正在执行任务,跳过本次触发: {groups}"
)
return False
else:
self._busy = True
# 激进模式: 先停止旧任务
if self.restart_mode == "aggressive":
await self._stop_bgi_gracefully()
async with self._lock:
self._busy = True
try:
cmd = [self.exe_path, "--startGroups"] + groups
self.logger.info(f"调用 BetterGI: {' '.join(cmd)}")
proc = await asyncio.create_subprocess_exec(
*cmd,
cwd=self.work_dir,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
stdout, stderr = await proc.communicate()
if proc.returncode in (0, 553):
self.logger.info(f"BetterGI 启动成功: {groups} (rc={proc.returncode})")
return True
else:
self.logger.error(
f"BetterGI 启动失败 rc={proc.returncode}: "
f"{stderr.decode('gbk', errors='replace')}"
)
return False
except Exception as e:
self.logger.exception(f"调用 BetterGI 异常: {e}")
return False
finally:
self._busy = False
# ============== 弹幕匹配引擎 ==============
class DanmuMatcher:
"""把弹幕文本匹配到对应规则。"""
def __init__(self, config: Config, cooldown: CooldownManager,
runner: BetterGIRunner, logger: logging.Logger):
self.config = config
self.cooldown = cooldown
self.runner = runner
self.logger = logger
async def handle(self, text: str, uid: int, uname: str):
text = (text or "").strip()
if not text:
return
admins = self.config.admin_uids
is_admin = uid in admins
for rule in self.config.rules:
if not self._match(text, rule):
continue
keyword = rule["keyword"]
groups = rule["groups"]
cooldown_sec = rule.get("cooldown", self.config.default_cooldown)
admin_only = rule.get("admin_only", False)
reply = rule.get("reply", "")
# 权限检查
if admin_only and not is_admin:
self.logger.info(
f"[权限不足] {uname}({uid}) 弹幕'{text}' 命中'{keyword}' 但需管理员"
)
return # 注意: 这里return了,继续检查下一条规则(如果有)
# 冷却检查
if not self.cooldown.can_fire(keyword, cooldown_sec):
remain = int(cooldown_sec - (time.time() - self.cooldown._last_fire.get(keyword, 0)))
self.logger.info(
f"[冷却中] '{keyword}' 还需 {remain}s, 来自 {uname}({uid})"
)
return # 冷却不触发,但日志已打印
# 触发
self.cooldown.mark_fired(keyword)
self.logger.info(
f"[触发] {uname}({uid}) 弹幕'{text}' -> 配置组 {groups}"
)
if reply:
self.logger.info(f"[回复提示] {reply}")
await self.runner.run_groups(groups)
return # 一条弹幕只触发第一个命中的规则
@staticmethod
def _match(text: str, rule: dict) -> bool:
keyword = rule["keyword"]
mtype = rule.get("match_type", "contains")
if mtype == "exact":
return text == keyword
elif mtype == "contains":
return keyword in text
elif mtype == "startswith":
return text.startswith(keyword)
elif mtype == "regex":
import re
return re.search(keyword, text) is not None
return False
# ============== B站弹幕协议 ==============
def make_packet(op: int, body: bytes = b"") -> bytes:
"""组装协议包。header 16字节 + body。"""
if isinstance(body, str):
body = body.encode("utf-8")
total = HEADER_LEN + len(body)
header = struct.pack(">IHHII", total, HEADER_LEN, 1, op, 1)
return header + body
def parse_packets(data: bytes):
"""一个 WebSocket 帧可能含多个协议包,循环切分。"""
offset = 0
packets = []
while offset < len(data):
if offset + HEADER_LEN > len(data):
break
total, header_len, proto_ver, op, seq = struct.unpack(
">IHHII", data[offset:offset + HEADER_LEN]
)
body = data[offset + header_len:offset + total]
packets.append((proto_ver, op, body))
offset += total
return packets
def decode_body(proto_ver: int, body: bytes) -> bytes:
"""按协议版本解压 body。"""
if proto_ver == PROTO_JSON:
return body
if proto_ver == PROTO_ZLIB:
return zlib.decompress(body)
if proto_ver == PROTO_BROTLI:
return brotli.decompress(body)
return body
def _clean_uid(value) -> int:
"""只接受纯十进制正整数 UID,匿名或含星号的值返回 0。"""
if isinstance(value, bool):
return 0
if isinstance(value, int):
return value if value > 0 else 0
if isinstance(value, str):
value = value.strip()
return int(value) if value.isdecimal() and int(value) > 0 else 0
return 0
def extract_danmu(body: bytes):
"""从消息体中提取弹幕(DANMU_MSG)。返回 [(text, uid, uname), ...]。"""
results = []
try:
decoded = decode_body(0, body) # body 已是解压后的,proto_ver 此处忽略
except Exception:
decoded = body
try:
msg = json.loads(decoded.decode("utf-8", errors="replace"))
except Exception:
return results
cmd = msg.get("cmd", "")
if cmd.startswith("DANMU_MSG"):
info = msg.get("info", [])
if not isinstance(info, list) or len(info) <= 2:
return results
text = str(info[1] if len(info) > 1 else "")
member = info[2]
if isinstance(member, (list, tuple)):
uid = _clean_uid(member[0] if len(member) > 0 else 0)
uname = str(member[1] if len(member) > 1 and member[1] is not None else "")
elif isinstance(member, dict):
uid = _clean_uid(member.get("uid_str") or member.get("uid") or member.get("mid"))
uname = str(member.get("uname") or member.get("name") or "")
else:
uid, uname = 0, ""
results.append((text, uid, uname))
return results
# ============== 弹幕客户端 ==============
class BliveClient:
"""B站直播弹幕 WebSocket 客户端,带自动重连。"""
def __init__(self, room_id: int, sessdata: str,
matcher: DanmuMatcher, logger: logging.Logger):
self.room_id = room_id
self.sessdata = sessdata
self.matcher = matcher
self.logger = logger
self._stop = False
self._buvid3 = ""
async def run(self):
"""主循环:断线自动重连,间隔递增。"""
# 启动时获取buvid3 (2024+协议认证必需)
self._buvid3 = get_buvid3(self.sessdata)
if self._buvid3:
self.logger.info(f"获取buvid3成功: {self._buvid3[:20]}...")
else:
self.logger.warning("获取buvid3失败,弹幕可能收不到! 将尝试匿名连接")
retry = 0
while not self._stop:
try:
# 每次连接前重新获取服务器地址(避免IP变化)
server_info = get_danmu_server(self.room_id)
ws_url = server_info["ws_url"]
token = server_info["token"]
self.logger.info(f"弹幕服务器: {server_info['host']}:{server_info['port']}")
await self._connect_once(ws_url, token)
retry = 0 # 连上后重置
except asyncio.CancelledError:
break
except Exception as e:
self.logger.warning(f"连接断开: {e}")
if self._stop:
break
retry += 1
wait = min(2 ** retry, 60) # 指数退避,最多60秒
self.logger.info(f"{wait}秒后重连(第{retry}次)...")
await asyncio.sleep(wait)
def stop(self):
self._stop = True
async def _connect_once(self, ws_url: str, token: str):
self.logger.info(f"连接直播间 room_id={self.room_id} ...")
# 禁用 websockets 自带 ping/pong (B站有自己的心跳协议)
async with websockets.connect(
ws_url,
max_size=None,
ping_interval=None,
ping_timeout=None,
close_timeout=5,
) as ws:
# 1. 发送认证包(使用API返回的token + buvid3)
auth = {
"uid": 0,
"roomid": self.room_id,
"protover": PROTO_BROTLI,
"platform": "web",
"type": 2,
"key": token or (self.sessdata or ""),
}
if self._buvid3:
auth["buvid"] = self._buvid3
await ws.send(make_packet(OP_AUTH, json.dumps(auth)))
self.logger.info("已发送认证包,等待响应...")
# 立即发送一个心跳,激活弹幕推送
await ws.send(make_packet(OP_HEARTBEAT))
# 2. 心跳任务
async def heartbeat():
while True:
await asyncio.sleep(15)
try:
await ws.send(make_packet(OP_HEARTBEAT))
except Exception:
break
hb_task = asyncio.create_task(heartbeat())
# 3. 接收循环
msg_counter = 0
async for raw in ws:
if isinstance(raw, str):
continue
for proto_ver, op, body in parse_packets(raw):
if op == OP_AUTH_REPLY:
# 检查认证是否真的成功
try:
auth_resp = json.loads(body.decode("utf-8", errors="replace"))
code = auth_resp.get("code", -1)
if code == 0:
self.logger.info(f"认证成功,开始监听弹幕 (响应: {auth_resp})")
else:
self.logger.error(f"认证失败! 响应: {auth_resp}")
except Exception:
self.logger.info(f"认证响应(原始): {body}")
elif op == OP_HEARTBEAT_REPLY:
# body 是 4 字节人气值(大端序int32)
if len(body) >= 4:
popularity = struct.unpack(">I", body[:4])[0]
self.logger.debug(f"人气值: {popularity}")
elif op == OP_MESSAGE:
# body 可能被压缩,按 proto_ver 解压后再切包
try:
decoded = decode_body(proto_ver, body)
except Exception as e:
self.logger.error(f"解压失败 proto={proto_ver} len={len(body)}: {e}")
continue
# 解压后可能内含多个子包
for sub_proto, sub_op, sub_body in parse_packets(decoded):
if sub_op == OP_MESSAGE:
msg_counter += 1
# 调试:打印每条消息的cmd类型
try:
msg_json = json.loads(sub_body.decode("utf-8", errors="replace"))
cmd = msg_json.get("cmd", "?")
self.logger.debug(f"#{msg_counter} cmd={cmd}")
except Exception:
self.logger.debug(f"#{msg_counter} 解析JSON失败")
continue
for text, uid, uname in extract_danmu(sub_body):
self.logger.info(f"[弹幕] {uname}({uid}): {text}")
await self.matcher.handle(text, uid, uname)
hb_task.cancel()
# ============== 日志 ==============
def setup_logger(config: Config) -> logging.Logger:
logger = logging.getLogger("danmu_bettergi")
logger.setLevel(getattr(logging, config.data["global"].get("log_level", "INFO")))
fmt = logging.Formatter("%(asctime)s [%(levelname)s] %(message)s",
datefmt="%H:%M:%S")
sh = logging.StreamHandler(sys.stdout)
sh.setFormatter(fmt)
logger.addHandler(sh)
log_file = config.data["global"].get("log_file")
if log_file:
LOG_DIR.mkdir(parents=True, exist_ok=True)
fh = logging.FileHandler(project_path(log_file, base=LOG_DIR), encoding="utf-8")
fh.setFormatter(fmt)
logger.addHandler(fh)
return logger
# ============== 入口 ==============
async def main():
ensure_runtime_dirs()
config_path = CONFIG_DIR / "config.json"
if not config_path.exists():
legacy_config_path = PROJECT_ROOT / "config.json"
if legacy_config_path.exists():
config_path = legacy_config_path
if not config_path.exists():
print("找不到 config.json,请先配置!")
sys.exit(1)
config = Config(str(config_path))
logger = setup_logger(config)
logger.info("=" * 50)
logger.info("BetterGI 弹幕联动启动")
logger.info(f"直播间: {config.room_id}")
logger.info(f"BetterGI: {config.bettergi_exe}")
logger.info(f"规则数: {len(config.rules)}")
logger.info(f"管理员UID: {config.admin_uids or ''}")
logger.info(f"默认冷却: {config.default_cooldown}s")
logger.info(f"重启模式: {config.restart_mode}")
logger.info("=" * 50)
cooldown = CooldownManager()
runner = BetterGIRunner(
config.bettergi_exe, config.bettergi_work_dir, logger,
restart_mode=config.restart_mode,
)
matcher = DanmuMatcher(config, cooldown, runner, logger)
client = BliveClient(config.room_id, config.sessdata, matcher, logger)
try:
await client.run()
except KeyboardInterrupt:
logger.info("收到退出信号,正在停止...")
client.stop()
if __name__ == "__main__":
try:
asyncio.run(main())
except KeyboardInterrupt:
pass
+10205
View File
File diff suppressed because it is too large Load Diff
+462
View File
@@ -0,0 +1,462 @@
from __future__ import annotations
import io
import multiprocessing
import os
import threading
import time
import traceback
import uuid
from multiprocessing.connection import Connection
from typing import Any, Callable
MAX_NEW_TOKENS = 384
DEFAULT_SYNTHESIS_TIMEOUT_SECONDS = 120.0
DEFAULT_STARTUP_TIMEOUT_SECONDS = 480.0
# 启动失败后再次拉起 worker 的最小间隔:避免"启动超时→立即重启→再超时"的死循环
# 在系统高负载时持续加载 torch/CUDA,进一步加剧卡顿。
STARTUP_FAILURE_BACKOFF_SECONDS = 60.0
DEFAULT_CPU_THREADS = 4
DEFAULT_CPU_AFFINITY_COUNT = 8
DEFAULT_PROCESS_PRIORITY = "below_normal"
class FasterQwenWorkerError(RuntimeError):
pass
class FasterQwenWorkerTimeout(FasterQwenWorkerError):
pass
def _generation_kwargs(settings: dict[str, Any], text: str) -> dict[str, Any]:
return {
"text": text,
"language": str(settings.get("language") or "Chinese"),
"non_streaming_mode": bool(settings.get("non_streaming_mode", True)),
"max_new_tokens": MAX_NEW_TOKENS,
}
def _bounded_int(value: Any, default: int, *, minimum: int, maximum: int) -> int:
try:
parsed = int(value)
except (TypeError, ValueError):
parsed = default
return max(minimum, min(maximum, parsed))
def _configure_worker_environment(settings: dict[str, Any]) -> int:
cpu_threads = _bounded_int(
settings.get("cpu_threads"),
DEFAULT_CPU_THREADS,
minimum=1,
maximum=8,
)
thread_value = str(cpu_threads)
for name in (
"OMP_NUM_THREADS",
"MKL_NUM_THREADS",
"OPENBLAS_NUM_THREADS",
"NUMEXPR_NUM_THREADS",
"VECLIB_MAXIMUM_THREADS",
"BLIS_NUM_THREADS",
):
os.environ[name] = thread_value
os.environ["TOKENIZERS_PARALLELISM"] = "false"
return cpu_threads
def _configure_torch_threads(torch_module: Any, cpu_threads: int) -> None:
torch_module.set_num_threads(cpu_threads)
try:
torch_module.set_num_interop_threads(1)
except RuntimeError:
# PyTorch only allows setting interop threads before parallel work starts.
pass
def _apply_worker_process_limits(settings: dict[str, Any]) -> dict[str, Any]:
cpu_count = max(1, int(os.cpu_count() or 1))
affinity_count = _bounded_int(
settings.get("cpu_affinity_count"),
DEFAULT_CPU_AFFINITY_COUNT,
minimum=0,
maximum=min(cpu_count, 63),
)
priority = str(settings.get("process_priority") or DEFAULT_PROCESS_PRIORITY).strip().lower()
applied_affinity = 0
applied_priority = "default"
if os.name == "nt":
try:
import ctypes
from ctypes import wintypes
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
kernel32.GetCurrentProcess.restype = wintypes.HANDLE
kernel32.SetPriorityClass.argtypes = [wintypes.HANDLE, wintypes.DWORD]
kernel32.SetPriorityClass.restype = wintypes.BOOL
kernel32.SetProcessAffinityMask.argtypes = [wintypes.HANDLE, ctypes.c_size_t]
kernel32.SetProcessAffinityMask.restype = wintypes.BOOL
process_handle = kernel32.GetCurrentProcess()
priority_classes = {
"idle": 0x00000040,
"below_normal": 0x00004000,
"normal": 0x00000020,
}
priority_class = priority_classes.get(priority, priority_classes[DEFAULT_PROCESS_PRIORITY])
if kernel32.SetPriorityClass(process_handle, priority_class):
applied_priority = priority if priority in priority_classes else DEFAULT_PROCESS_PRIORITY
if affinity_count > 0:
affinity_mask = (1 << affinity_count) - 1
if kernel32.SetProcessAffinityMask(process_handle, affinity_mask):
applied_affinity = affinity_count
except Exception:
pass
return {
"cpu_affinity_count": applied_affinity,
"process_priority": applied_priority,
}
def _load_runtime(settings: dict[str, Any]) -> dict[str, Any]:
device = str(settings.get("device") or "cuda")
if device == "cpu" and "CUDA_VISIBLE_DEVICES" not in os.environ:
os.environ["CUDA_VISIBLE_DEVICES"] = ""
cpu_threads = _configure_worker_environment(settings)
import torch
_configure_torch_threads(torch, cpu_threads)
from faster_qwen3_tts import FasterQwen3TTS
if device == "cpu":
torch.cuda.is_available = lambda: False
load_kwargs: dict[str, Any] = {}
if device == "cpu":
load_kwargs["device"] = "cpu"
model = FasterQwen3TTS.from_pretrained(
str(settings.get("model_name_or_path") or "Qwen/Qwen3-TTS-12Hz-0.6B-Base"),
**load_kwargs,
)
# voice_clone_prompt 不再预计算:当前 faster_qwen3_tts 的 FasterQwen3TTS 没有
# create_voice_clone_prompt 方法,预计算只会失败。改为每次合成时在
# _synthesize_wav 里直接传 ref_audio/ref_text/xvec_only 参数。
voice_clone_prompt = None
return {
"model": model,
"torch": torch,
"voice_clone_prompt": voice_clone_prompt,
"settings": settings,
"cpu_threads": cpu_threads,
}
def _synthesize_wav(runtime: dict[str, Any], text: str) -> bytes:
import soundfile as sf
model = runtime["model"]
torch = runtime["torch"]
settings = runtime["settings"]
safe_text = str(text or "").strip()[:80] or "欢迎来到直播间。"
kwargs = _generation_kwargs(settings, safe_text)
voice_clone_prompt = runtime.get("voice_clone_prompt")
if voice_clone_prompt is not None:
kwargs["voice_clone_prompt"] = voice_clone_prompt
else:
ref_audio = str(settings.get("ref_audio") or "")
if not ref_audio:
raise RuntimeError("Faster-Qwen3-TTS requires ref_audio")
kwargs.update({
"ref_audio": ref_audio,
"ref_text": str(settings.get("ref_text") or "") or None,
"xvec_only": bool(settings.get("xvec_only", True)),
"append_silence": bool(settings.get("append_silence", True)),
})
if not torch.cuda.is_available():
torch.backends.cudnn.enabled = False
with torch.inference_mode():
wavs, sample_rate = model.generate_voice_clone(**kwargs)
output = io.BytesIO()
audio = wavs[0]
if isinstance(audio, torch.Tensor):
audio = audio.cpu().numpy()
sf.write(output, audio, sample_rate, format="WAV")
return output.getvalue()
def faster_qwen_worker_main(connection: Connection, settings: dict[str, Any]) -> None:
try:
cpu_threads = _configure_worker_environment(settings)
process_limits = _apply_worker_process_limits(settings)
load_started = time.monotonic()
runtime = _load_runtime(settings)
load_ms = int((time.monotonic() - load_started) * 1000)
warmup_started = time.monotonic()
_synthesize_wav(runtime, "系统启动")
warmup_ms = int((time.monotonic() - warmup_started) * 1000)
connection.send({
"type": "ready",
"pid": os.getpid(),
"load_ms": load_ms,
"warmup_ms": warmup_ms,
"max_new_tokens": MAX_NEW_TOKENS,
"cpu_threads": cpu_threads,
**process_limits,
})
except BaseException as exc:
try:
connection.send({
"type": "startup_error",
"error_type": type(exc).__name__,
"error": str(exc),
"traceback": traceback.format_exc(),
})
except Exception:
pass
connection.close()
return
while True:
try:
message = connection.recv()
except (EOFError, OSError):
break
if not isinstance(message, dict):
continue
command = str(message.get("command") or "")
if command == "stop":
break
if command != "synthesize":
continue
request_id = str(message.get("request_id") or "")
started = time.monotonic()
try:
audio = _synthesize_wav(runtime, str(message.get("text") or ""))
connection.send({
"type": "result",
"request_id": request_id,
"audio": audio,
"duration_ms": int((time.monotonic() - started) * 1000),
"bytes": len(audio),
"max_new_tokens": MAX_NEW_TOKENS,
})
except BaseException as exc:
try:
connection.send({
"type": "error",
"request_id": request_id,
"error_type": type(exc).__name__,
"error": str(exc),
"traceback": traceback.format_exc(),
})
except Exception:
break
connection.close()
class FasterQwenWorkerClient:
def __init__(
self,
settings: dict[str, Any],
logger=None,
*,
synthesis_timeout_seconds: float = DEFAULT_SYNTHESIS_TIMEOUT_SECONDS,
startup_timeout_seconds: float = DEFAULT_STARTUP_TIMEOUT_SECONDS,
context=None,
process_target: Callable[..., None] | None = None,
):
self.settings = dict(settings)
self.logger = logger
self.synthesis_timeout_seconds = max(1.0, float(synthesis_timeout_seconds))
self.startup_timeout_seconds = max(10.0, float(startup_timeout_seconds))
self._context = context or multiprocessing.get_context("spawn")
self._process_target = process_target or faster_qwen_worker_main
self._lock = threading.RLock()
self._process = None
self._connection = None
self._worker_pid = 0
self._next_start_after = 0.0
@property
def worker_pid(self) -> int:
return int(self._worker_pid or 0)
def _log(self, level: str, message: str, *args) -> None:
if self.logger:
getattr(self.logger, level)(message, *args)
def _is_alive_locked(self) -> bool:
return bool(self._process is not None and self._process.is_alive())
def ensure_ready(self) -> dict[str, Any]:
with self._lock:
if self._is_alive_locked() and self._connection is not None:
return {"pid": self.worker_pid, "reused": True}
return self._start_worker_locked()
def _start_worker_locked(self) -> dict[str, Any]:
now = time.monotonic()
if now < self._next_start_after:
wait = int(self._next_start_after - now)
raise FasterQwenWorkerError(
f"TTS worker 启动退避中,距上次启动失败不足 {int(STARTUP_FAILURE_BACKOFF_SECONDS)} 秒,"
f"{wait} 秒后可重试"
)
self._terminate_worker_locked("replace_stale_worker", graceful=False)
parent_connection, child_connection = self._context.Pipe(duplex=True)
process = self._context.Process(
target=self._process_target,
args=(child_connection, self.settings),
name="FasterQwen3TTSWorker",
daemon=True,
)
process.start()
try:
child_connection.close()
except Exception:
pass
self._process = process
self._connection = parent_connection
self._worker_pid = int(getattr(process, "pid", 0) or 0)
self._log("info", "[FasterQwenTTS] worker 已启动, pid=%s,正在加载和预热", self.worker_pid)
if not parent_connection.poll(self.startup_timeout_seconds):
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
self._terminate_worker_locked("startup_timeout", graceful=False)
raise FasterQwenWorkerError(
f"Faster-Qwen3-TTS worker startup exceeded {self.startup_timeout_seconds:.0f}s"
)
try:
message = parent_connection.recv()
except (EOFError, OSError) as exc:
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
self._terminate_worker_locked("startup_connection_closed", graceful=False)
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker exited during startup") from exc
if not isinstance(message, dict) or message.get("type") != "ready":
if isinstance(message, dict):
error = str(message.get("error") or message.get("error_type") or "unknown startup error")
else:
error = "invalid startup response"
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
self._terminate_worker_locked("startup_error", graceful=False)
raise FasterQwenWorkerError(error)
self._next_start_after = 0.0
self._log(
"info",
"[FasterQwenTTS] worker 预热完成, pid=%s load=%sms warmup=%sms max_new_tokens=%s",
self.worker_pid,
message.get("load_ms"),
message.get("warmup_ms"),
message.get("max_new_tokens"),
)
self._log(
"info",
"[FasterQwenTTS] worker 资源限制: cpu_threads=%s affinity=%s priority=%s",
message.get("cpu_threads"),
message.get("cpu_affinity_count"),
message.get("process_priority"),
)
return message
def synthesize(self, text: str) -> tuple[bytes, dict[str, Any]]:
with self._lock:
self.ensure_ready()
request_id = uuid.uuid4().hex
connection = self._connection
try:
connection.send({
"command": "synthesize",
"request_id": request_id,
"text": str(text or ""),
})
except (BrokenPipeError, EOFError, OSError) as exc:
self._restart_after_failure_locked("send_failed")
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker connection failed") from exc
if not connection.poll(self.synthesis_timeout_seconds):
self._log(
"error",
"[FasterQwenTTS] 单次合成超过 %.0f 秒,强制终止 worker pid=%s",
self.synthesis_timeout_seconds,
self.worker_pid,
)
restart_error = self._restart_after_failure_locked("synthesis_timeout")
suffix = f"; restart failed: {restart_error}" if restart_error else ""
raise FasterQwenWorkerTimeout(
f"Faster-Qwen3-TTS synthesis exceeded {self.synthesis_timeout_seconds:.0f}s{suffix}"
)
try:
message = connection.recv()
except (EOFError, OSError) as exc:
self._restart_after_failure_locked("worker_exited")
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker exited during synthesis") from exc
if not isinstance(message, dict) or message.get("request_id") != request_id:
self._restart_after_failure_locked("invalid_response")
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker returned an invalid response")
if message.get("type") == "error":
error = str(message.get("error") or message.get("error_type") or "synthesis failed")
lowered = error.lower()
if "cuda" in lowered or "out of memory" in lowered or "device-side" in lowered:
self._restart_after_failure_locked("cuda_error")
raise FasterQwenWorkerError(error)
if message.get("type") != "result" or not isinstance(message.get("audio"), bytes):
self._restart_after_failure_locked("invalid_result")
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker returned no audio")
return message["audio"], message
def _restart_after_failure_locked(self, reason: str) -> str:
self._terminate_worker_locked(reason, graceful=False)
try:
self._start_worker_locked()
return ""
except Exception as exc:
self._log("error", "[FasterQwenTTS] worker 自动重启失败: %s", exc)
return str(exc)
def _terminate_worker_locked(self, reason: str, *, graceful: bool) -> None:
process = self._process
connection = self._connection
self._process = None
self._connection = None
self._worker_pid = 0
if process is None:
if connection is not None:
try:
connection.close()
except Exception:
pass
return
if graceful and process.is_alive() and connection is not None:
try:
connection.send({"command": "stop"})
process.join(timeout=3.0)
except Exception:
pass
if process.is_alive():
self._log("warning", "[FasterQwenTTS] 终止 worker, reason=%s pid=%s", reason, process.pid)
process.terminate()
process.join(timeout=10.0)
if process.is_alive():
process.kill()
process.join(timeout=5.0)
if connection is not None:
try:
connection.close()
except Exception:
pass
def close(self) -> None:
with self._lock:
self._terminate_worker_locked("shutdown", graceful=True)
+569
View File
@@ -0,0 +1,569 @@
"""安全、幂等地将旧版 JSON/日志快照补录到统计数据库。"""
from __future__ import annotations
import argparse
import asyncio
import hashlib
import json
import re
from collections import Counter
from collections.abc import Mapping
from datetime import UTC, datetime, timedelta, timezone
from pathlib import Path
from typing import Any
try:
from .stats_store import DEFAULT_DATABASE_PATH, StatsStore, business_date
except ImportError: # 允许直接执行文件
from stats_store import DEFAULT_DATABASE_PATH, StatsStore, business_date
BEIJING_TZ = timezone(timedelta(hours=8), name="Asia/Shanghai")
_SOURCE_FILES = (
"users.json",
"song_requests.json",
"admin_audit.log",
"queue_state.json",
"music_state.json",
"tts_state.json",
)
_USER_FIELDS = {
"uname",
"points",
"last_signin_date",
"created_at",
"blocked_all",
"blocked_queue",
"blocked_song_request",
"note",
}
_SONG_FIELDS = {
"id",
"name",
"artist",
"duration_ms",
"duration_sec",
"keyword",
"uid",
"uname",
"requested_at",
"source",
"remove_after_play",
"started_at",
"finished_at",
"status",
}
_AUDIT_FIELDS = {"at", "action", "target", "client_ip", "session_id", "detail"}
TABLE_COLUMNS = {
"users": {
"platform", "platform_user_id", "display_name", "avatar_url", "user_level",
"is_admin", "first_seen_at_utc", "last_seen_at_utc", "snapshot_json",
},
"song_requests": {
"request_id", "platform", "platform_user_id", "requested_at_utc", "business_date",
"song_id", "song_name", "artist", "source", "status", "points_cost",
"queue_position", "payload_json",
},
"playback_sessions": {
"playback_id", "request_id", "song_id", "song_name", "started_at_utc",
"ended_at_utc", "business_date", "status", "duration_ms", "played_ms",
"stop_reason", "payload_json",
},
"admin_audit_events": {
"event_id", "occurred_at_utc", "business_date", "actor", "action",
"target_type", "target_id", "success", "remote_address_hash", "payload_json",
},
"point_transactions": {
"transaction_id", "platform", "platform_user_id", "occurred_at_utc",
"business_date", "amount", "balance_after", "reason", "reference_type",
"reference_id", "payload_json",
},
"events": {
"event_id", "event_type", "category", "occurred_at_utc", "business_date",
"payload_json",
},
}
_POINT_DETAIL_RE = re.compile(r"^delta=(-?\d+)\s+now=(-?\d+)$")
_COUNT_DETAIL_RE = re.compile(r"^count=(\d+)$")
_STATUS_MAP = {
"played": "completed",
"complete": "completed",
"finished": "completed",
"success": "completed",
"skipped": "skipped",
"skip": "skipped",
"interrupted": "interrupted",
"cancelled": "cancelled",
"canceled": "cancelled",
"play_error": "failed",
"error": "failed",
"failed": "failed",
"queued": "queued",
"pending": "queued",
"playing": "playing",
"active": "playing",
}
def _hash(*parts: Any) -> str:
encoded = json.dumps(parts, ensure_ascii=False, separators=(",", ":"), default=str).encode("utf-8", "replace")
return hashlib.sha256(encoded).hexdigest()
def _text(value: Any, limit: int) -> str | None:
if value is None:
return None
result = str(value).strip()
return result[:limit] if result else None
def _int(value: Any, *, minimum: int | None = None) -> int | None:
try:
result = int(value)
except (TypeError, ValueError, OverflowError):
return None
if minimum is not None and result < minimum:
return None
return result
def _bool(value: Any) -> bool:
return value is True or value == 1
def _utc_iso(
value: Any,
*,
naive_is_beijing: bool = True,
fallback: datetime | None = None,
timespec: str = "milliseconds",
) -> str:
parsed: datetime
if isinstance(value, (int, float)) and not isinstance(value, bool):
parsed = datetime.fromtimestamp(float(value), UTC)
elif isinstance(value, str) and value.strip():
raw = value.strip()
try:
parsed = datetime.fromisoformat(raw.replace("Z", "+00:00"))
except ValueError:
parsed = fallback or datetime.now(UTC)
else:
parsed = fallback or datetime.now(UTC)
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=BEIJING_TZ if naive_is_beijing else UTC)
return parsed.astimezone(UTC).isoformat(timespec=timespec).replace("+00:00", "Z")
def _file_time(path: Path) -> str:
return _utc_iso(datetime.fromtimestamp(path.stat().st_mtime, UTC), naive_is_beijing=False)
def _read_bytes(path: Path) -> bytes:
return path.read_bytes()
def _read_json(raw: bytes) -> Any:
return json.loads(raw.decode("utf-8-sig"))
def _fingerprint(raw: bytes) -> str:
return hashlib.sha256(raw).hexdigest()
def _normalize_status(value: Any, default: str) -> str:
key = str(value or "").strip().casefold()
return _STATUS_MAP.get(key, default)
def _json_payload(value: Mapping[str, Any]) -> dict[str, Any]:
return {key: item for key, item in value.items() if item is not None}
def _user_records(path: Path, raw: bytes) -> list[tuple[str, dict[str, Any]]]:
data = _read_json(raw)
if not isinstance(data, Mapping):
return []
snapshot_at = _file_time(path)
records: list[tuple[str, dict[str, Any]]] = []
for uid, original in data.items():
if not isinstance(original, Mapping):
continue
item = {key: original[key] for key in _USER_FIELDS if key in original}
platform_user_id = _text(uid, 64)
if not platform_user_id:
continue
created_at = _utc_iso(item.get("created_at"), fallback=datetime.fromtimestamp(path.stat().st_mtime, UTC))
snapshot = _json_payload({
"points": _int(item.get("points")),
"last_signin_date": _text(item.get("last_signin_date"), 10),
"blocked_all": _bool(item.get("blocked_all")),
"blocked_queue": _bool(item.get("blocked_queue")),
"blocked_song_request": _bool(item.get("blocked_song_request")),
"note_length": len(str(item.get("note") or "")),
"legacy_snapshot": True,
})
records.append(("users", {
"platform": "bilibili",
"platform_user_id": platform_user_id,
"display_name": _text(item.get("uname"), 128),
"avatar_url": None,
"user_level": None,
"is_admin": False,
"first_seen_at_utc": created_at,
"last_seen_at_utc": snapshot_at,
"snapshot_json": snapshot,
}))
return records
def _song_records(path: Path, raw: bytes) -> list[tuple[str, dict[str, Any]]]:
data = _read_json(raw)
if not isinstance(data, Mapping):
return []
snapshot_at = _file_time(path)
fallback = datetime.fromtimestamp(path.stat().st_mtime, UTC)
entries: list[dict[str, Any]] = []
sections = (("queue", data.get("queue")), ("active", [data.get("active")]), ("history", data.get("history")))
for section, values in sections:
if not isinstance(values, list):
continue
section_entries: dict[tuple[str | None, str, str | None, str], list[dict[str, Any]]] = {}
for queue_position, original in enumerate(values, 1):
if not isinstance(original, Mapping):
continue
item = {key: original[key] for key in _SONG_FIELDS if key in original}
requested = _utc_iso(item.get("requested_at"), fallback=fallback, timespec="microseconds")
source = _text(item.get("source"), 64) or "legacy"
business_key = (
_text(item.get("uid"), 64),
requested,
_text(item.get("id"), 128),
source,
)
entry = {
"section": section,
"item": item,
"requested": requested,
"started": _utc_iso(item["started_at"], timespec="microseconds") if item.get("started_at") is not None else None,
"ended": _utc_iso(item["finished_at"], timespec="microseconds") if item.get("finished_at") is not None else None,
"source": source,
"business_key": business_key,
"queue_position": queue_position if section == "queue" else None,
}
section_entries.setdefault(business_key, []).append(entry)
for business_key, duplicates in section_entries.items():
duplicates.sort(key=lambda entry: (
entry["started"] or "",
entry["ended"] or "",
str(entry["item"].get("status") or ""),
str(entry["item"].get("name") or ""),
str(entry["item"].get("artist") or ""),
str(entry["item"].get("duration_ms") or ""),
))
for duplicate_ordinal, entry in enumerate(duplicates):
entry["request_id"] = _hash("legacy-song-request", *business_key, duplicate_ordinal)
entries.append(entry)
section_priority = {"queue": 0, "active": 1, "history": 2}
requests: dict[str, dict[str, Any]] = {}
playbacks: dict[str, dict[str, Any]] = {}
for entry in sorted(entries, key=lambda value: section_priority[value["section"]]):
section = entry["section"]
item = entry["item"]
request_id = entry["request_id"]
confidence = "snapshot" if section in {"queue", "active"} else "history"
default_status = "queued" if section == "queue" else "playing" if section == "active" else "unknown"
status = _normalize_status(item.get("status"), default_status)
requests[request_id] = {
"request_id": request_id,
"platform": "bilibili",
"platform_user_id": _text(item.get("uid"), 64),
"requested_at_utc": entry["requested"],
"business_date": business_date(entry["requested"]),
"song_id": _text(item.get("id"), 128),
"song_name": _text(item.get("name"), 256),
"artist": _text(item.get("artist"), 256),
"source": entry["source"],
"status": status,
"points_cost": 0,
"queue_position": entry["queue_position"],
"payload_json": _json_payload({
"legacy_section": section,
"confidence": confidence,
"latest_80_only": section == "history",
"keyword_length": len(str(item.get("keyword") or "")),
"remove_after_play": _bool(item.get("remove_after_play")),
"snapshot_at_utc": snapshot_at if confidence == "snapshot" else None,
}),
}
started = entry["started"]
if not started:
continue
ended = entry["ended"]
played_ms = None
if ended:
start_dt = datetime.fromisoformat(started.replace("Z", "+00:00"))
end_dt = datetime.fromisoformat(ended.replace("Z", "+00:00"))
played_ms = max(0, round((end_dt - start_dt).total_seconds() * 1000))
playback_id = _hash("legacy-playback", request_id, started)
playbacks[playback_id] = {
"playback_id": playback_id,
"request_id": request_id,
"song_id": _text(item.get("id"), 128),
"song_name": _text(item.get("name"), 256),
"started_at_utc": started,
"ended_at_utc": ended,
"business_date": business_date(started),
"status": status,
"duration_ms": _int(item.get("duration_ms"), minimum=0),
"played_ms": played_ms,
"stop_reason": "snapshot" if confidence == "snapshot" else status,
"payload_json": {
"legacy_section": section,
"confidence": confidence,
"latest_80_only": section == "history",
},
}
return [
*(("song_requests", record) for record in requests.values()),
*(("playback_sessions", record) for record in playbacks.values()),
]
def _safe_detail(detail: Any) -> dict[str, Any]:
text = str(detail or "")
result: dict[str, Any] = {
"detail_length": len(text),
"detail_sha256": _hash("admin-detail", text),
}
count_match = _COUNT_DETAIL_RE.fullmatch(text)
if count_match:
result.update({"detail_kind": "count", "count": int(count_match.group(1))})
return result
try:
parsed = json.loads(text)
except (TypeError, ValueError):
parsed = None
if isinstance(parsed, Mapping):
allowed_flags = {key: _bool(parsed[key]) for key in ("blocked_all", "blocked_queue", "blocked_song_request") if key in parsed}
if allowed_flags:
result.update({"detail_kind": "flags", "flags": allowed_flags})
return result
result["detail_kind"] = "opaque"
return result
def _audit_records(raw: bytes) -> list[tuple[str, dict[str, Any]]]:
records: list[tuple[str, dict[str, Any]]] = []
duplicate_counts: Counter[str] = Counter()
for raw_line in raw.decode("utf-8-sig", "replace").splitlines():
line = raw_line.strip()
if not line:
continue
try:
original = json.loads(line)
except json.JSONDecodeError:
continue
if not isinstance(original, Mapping):
continue
line_hash = _hash("legacy-admin-audit-line", line)
duplicate_ordinal = duplicate_counts[line_hash]
duplicate_counts[line_hash] += 1
item = {key: original[key] for key in _AUDIT_FIELDS if key in original}
occurred = _utc_iso(item.get("at"), naive_is_beijing=True)
action = _text(item.get("action"), 96) or "unknown"
target = _text(item.get("target"), 256)
event_id = _hash("legacy-admin-audit", line_hash, duplicate_ordinal)
payload = _safe_detail(item.get("detail"))
session = _text(item.get("session_id"), 512)
if session:
payload["session_hash"] = _hash("admin-session", session)
records.append(("admin_audit_events", {
"event_id": event_id,
"occurred_at_utc": occurred,
"business_date": business_date(occurred),
"actor": "admin",
"action": action,
"target_type": "legacy_target",
"target_id": target,
"success": True,
"remote_address_hash": _hash("admin-ip", item.get("client_ip")) if item.get("client_ip") else None,
"payload_json": payload,
}))
point_match = _POINT_DETAIL_RE.fullmatch(str(item.get("detail") or "")) if action == "user_add_points" else None
if point_match and target:
records.append(("point_transactions", {
"transaction_id": _hash("legacy-admin-points", event_id),
"platform": "bilibili",
"platform_user_id": target,
"occurred_at_utc": occurred,
"business_date": business_date(occurred),
"amount": int(point_match.group(1)),
"balance_after": int(point_match.group(2)),
"reason": "admin_adjustment",
"reference_type": "admin_audit",
"reference_id": event_id,
"payload_json": {"legacy_import": True},
}))
return records
def _snapshot_event(path: Path, raw: bytes) -> list[tuple[str, dict[str, Any]]]:
data = _read_json(raw)
if not isinstance(data, Mapping):
return []
monitor_updated_at = (data.get("monitor") or {}).get("updated_at") if isinstance(data.get("monitor"), Mapping) else None
occurred = _utc_iso(
data.get("updated_at") or monitor_updated_at,
fallback=datetime.fromtimestamp(path.stat().st_mtime, UTC),
)
if path.name == "queue_state.json":
queue = data.get("queue") if isinstance(data.get("queue"), list) else []
payload = {
"queue_size": len(queue),
"has_active_user": bool(data.get("current_admin_uid")),
"has_group": bool(data.get("current_group")),
"default_running": _bool(data.get("default_running")),
"login_status": _text(data.get("login_status"), 64),
"has_user_finished_once": _bool(data.get("has_user_finished_once")),
"confidence": "snapshot",
}
event_type = "legacy.queue_snapshot"
category = "queue"
elif path.name == "music_state.json":
current = data.get("current") if isinstance(data.get("current"), Mapping) else {}
monitor = data.get("monitor") if isinstance(data.get("monitor"), Mapping) else {}
payload = {
"playing": _bool(data.get("playing")),
"current_title": _text(current.get("title"), 256),
"current_artist": _text(current.get("artist"), 256),
"duration": _int(current.get("duration"), minimum=0),
"progress": _int(current.get("progress"), minimum=0),
"playlist_size": len(data.get("playlist")) if isinstance(data.get("playlist"), list) else 0,
"request_size": len(data.get("requests")) if isinstance(data.get("requests"), list) else 0,
"monitor_online": _bool(monitor.get("online")),
"platform": _text(monitor.get("platform"), 32),
"confidence": "snapshot",
}
event_type = "legacy.music_snapshot"
category = "music"
else:
recent = data.get("recent_events") if isinstance(data.get("recent_events"), list) else []
payload = {
"enabled": _bool(data.get("enabled")),
"provider": _text(data.get("provider"), 64),
"model_loaded": _bool(data.get("model_loaded")),
"last_duration_ms": _int(data.get("last_duration_ms"), minimum=0),
"has_last_error": bool(data.get("last_error")),
"total_synthesized": _int(data.get("total_synthesized"), minimum=0),
"total_errors": _int(data.get("total_errors"), minimum=0),
"recent_event_count": len(recent),
"last_text_length": len(str(data.get("last_text") or "")),
"confidence": "snapshot",
}
event_type = "legacy.tts_snapshot"
category = "tts"
payload = _json_payload(payload)
return [("events", {
"event_id": _hash("legacy-snapshot", path.name, _fingerprint(raw)),
"event_type": event_type,
"category": category,
"occurred_at_utc": occurred,
"business_date": business_date(occurred),
"payload_json": payload,
})]
def _build_records(path: Path, raw: bytes) -> list[tuple[str, dict[str, Any]]]:
if path.name == "users.json":
return _user_records(path, raw)
if path.name == "song_requests.json":
return _song_records(path, raw)
if path.name == "admin_audit.log":
return _audit_records(raw)
return _snapshot_event(path, raw)
def validate_records(records: list[tuple[str, dict[str, Any]]]) -> None:
"""拒绝未知表、未知字段及非映射记录,避免导入边界被意外扩大。"""
for record_number, record in enumerate(records, 1):
if not isinstance(record, tuple) or len(record) != 2:
raise ValueError(f"{record_number} 条导入记录格式无效")
table, fields = record
allowed_columns = TABLE_COLUMNS.get(table)
if allowed_columns is None:
raise ValueError(f"{record_number} 条记录使用未知表: {table}")
if not isinstance(fields, Mapping):
raise ValueError(f"{record_number} 条记录字段不是映射")
unknown_columns = set(fields) - allowed_columns
if unknown_columns:
raise ValueError(f"{table} 包含未知字段: {sorted(unknown_columns)}")
if not fields:
raise ValueError(f"{table} 导入记录不能为空")
async def backfill_legacy_statistics(store: StatsStore, data_dir: str | Path, dry_run: bool = False) -> dict[str, Any]:
"""扫描旧数据并通过 ``StatsStore.import_once`` 原子、幂等地补录。"""
root = Path(data_dir)
result: dict[str, Any] = {"dry_run": bool(dry_run), "sources": {}, "record_counts": {}}
totals: Counter[str] = Counter()
for filename in _SOURCE_FILES:
path = root / filename
if not path.is_file():
result["sources"][filename] = {"status": "missing", "records": 0}
continue
raw = _read_bytes(path)
records = _build_records(path, raw)
validate_records(records)
counts = Counter(table for table, _ in records)
totals.update(counts)
imported = True
if not dry_run:
imported = await store.import_once(
f"legacy-history:{filename}",
"legacy_history",
records,
fingerprint=_fingerprint(raw),
checkpoint_key="sha256",
checkpoint_value=_fingerprint(raw),
metadata={"filename": filename, "record_count": len(records), "schema": 1},
)
result["sources"][filename] = {
"status": "dry_run" if dry_run else "completed" if imported else "failed",
"records": len(records),
"tables": dict(sorted(counts.items())),
}
result["record_counts"] = dict(sorted(totals.items()))
result["total_records"] = sum(totals.values())
return result
async def _main_async(args: argparse.Namespace) -> int:
data_dir = Path(args.data_dir)
if args.dry_run:
result = await backfill_legacy_statistics(StatsStore(args.database), data_dir, dry_run=True)
else:
store = StatsStore(args.database)
if not await store.start():
print(json.dumps({"status": "failed", "reason": "database_start_failed"}, ensure_ascii=False))
return 1
try:
result = await backfill_legacy_statistics(store, data_dir)
await store.flush()
finally:
await store.close()
print(json.dumps(result, ensure_ascii=False, sort_keys=True))
return 0 if all(source["status"] != "failed" for source in result["sources"].values()) else 1
def main() -> int:
parser = argparse.ArgumentParser(description="安全、幂等地补录旧版直播统计")
parser.add_argument("--data-dir", default=str(Path(__file__).resolve().parents[1] / "data"))
parser.add_argument("--database", default=str(DEFAULT_DATABASE_PATH))
parser.add_argument("--dry-run", action="store_true")
return asyncio.run(_main_async(parser.parse_args()))
if __name__ == "__main__":
raise SystemExit(main())
+260
View File
@@ -0,0 +1,260 @@
"""Unified launcher for source and frozen builds."""
from __future__ import annotations
import argparse
import asyncio
import ctypes
import os
import socket
import subprocess
import sys
from pathlib import Path
APP_DIR = Path(__file__).resolve().parent
if str(APP_DIR) not in sys.path:
sys.path.insert(0, str(APP_DIR))
from core.runtime_paths import APP_ROOT, DATA_DIR, ensure_runtime_dirs
ERROR_ALREADY_EXISTS = 183
MAIN_INSTANCE_MUTEX = "Local\\BetterGI_LiveStreaming_Main_5191"
class SingleInstanceLock:
"""Windows named mutex used by the queue-producing main process only."""
def __init__(self, name: str = MAIN_INSTANCE_MUTEX):
self.name = name
self._handle = None
def acquire(self) -> bool:
if os.name != "nt":
return True
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
handle = kernel32.CreateMutexW(None, False, self.name)
last_error = ctypes.get_last_error()
if not handle:
raise ctypes.WinError(last_error)
self._handle = handle
if last_error == ERROR_ALREADY_EXISTS:
self.close()
return False
return True
def close(self):
if self._handle is None or os.name != "nt":
return
ctypes.WinDLL("kernel32", use_last_error=True).CloseHandle(self._handle)
self._handle = None
class _JOBOBJECT_IO_COUNTERS(ctypes.Structure):
_fields_ = [
("ReadOperationCount", ctypes.c_ulonglong),
("WriteOperationCount", ctypes.c_ulonglong),
("OtherOperationCount", ctypes.c_ulonglong),
("ReadTransferCount", ctypes.c_ulonglong),
("WriteTransferCount", ctypes.c_ulonglong),
("OtherTransferCount", ctypes.c_ulonglong),
]
class _JOBOBJECT_BASIC_LIMIT_INFORMATION(ctypes.Structure):
_fields_ = [
("PerProcessUserTimeLimit", ctypes.c_longlong),
("PerJobUserTimeLimit", ctypes.c_longlong),
("LimitFlags", ctypes.c_ulong),
("MinimumWorkingSetSize", ctypes.c_size_t),
("MaximumWorkingSetSize", ctypes.c_size_t),
("ActiveProcessLimit", ctypes.c_ulong),
("Affinity", ctypes.c_size_t),
("PriorityClass", ctypes.c_ulong),
("SchedulingClass", ctypes.c_ulong),
]
class _JOBOBJECT_EXTENDED_LIMIT_INFORMATION(ctypes.Structure):
_fields_ = [
("BasicLimitInformation", _JOBOBJECT_BASIC_LIMIT_INFORMATION),
("IoInfo", _JOBOBJECT_IO_COUNTERS),
("ProcessMemoryLimit", ctypes.c_size_t),
("JobMemoryLimit", ctypes.c_size_t),
("PeakProcessMemoryUsed", ctypes.c_size_t),
("PeakJobMemoryUsed", ctypes.c_size_t),
]
class WindowsJob:
"""Kill spawned music/TTS processes automatically when the launcher exits."""
JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE = 0x00002000
JOB_OBJECT_EXTENDED_LIMIT_INFORMATION = 9
def __init__(self):
self._handle = None
if os.name != "nt":
return
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
handle = kernel32.CreateJobObjectW(None, None)
if not handle:
raise ctypes.WinError(ctypes.get_last_error())
info = _JOBOBJECT_EXTENDED_LIMIT_INFORMATION()
info.BasicLimitInformation.LimitFlags = self.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
ok = kernel32.SetInformationJobObject(
handle,
self.JOB_OBJECT_EXTENDED_LIMIT_INFORMATION,
ctypes.byref(info),
ctypes.sizeof(info),
)
if not ok:
error = ctypes.get_last_error()
kernel32.CloseHandle(handle)
raise ctypes.WinError(error)
self._handle = handle
def assign(self, process: subprocess.Popen):
if self._handle is None or os.name != "nt":
return
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
if not kernel32.AssignProcessToJobObject(self._handle, process._handle):
raise ctypes.WinError(ctypes.get_last_error())
def close(self):
if self._handle is None or os.name != "nt":
return
ctypes.WinDLL("kernel32", use_last_error=True).CloseHandle(self._handle)
self._handle = None
def _is_frozen() -> bool:
return bool(getattr(sys, "frozen", False))
def _role_command(role: str, port: int, host: str) -> list[str]:
if _is_frozen():
return [sys.executable, "--role", role, "--port", str(port), "--host", host]
return [sys.executable, str(Path(__file__).resolve()), "--role", role, "--port", str(port), "--host", host]
def _assert_port_available(host: str, port: int):
probe_host = "0.0.0.0" if host in {"", "::"} else host
family = socket.AF_INET6 if ":" in probe_host else socket.AF_INET
with socket.socket(family, socket.SOCK_STREAM) as sock:
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 0)
sock.bind((probe_host, port))
def _spawn_role(role: str, port: int, host: str, *, visible: bool = False) -> subprocess.Popen:
creationflags = 0
if os.name == "nt":
creationflags = 0x00000010 if visible else 0x08000000 # CREATE_NEW_CONSOLE / CREATE_NO_WINDOW
return subprocess.Popen(
_role_command(role, port, host),
cwd=APP_ROOT,
stdout=None if visible else subprocess.DEVNULL,
stderr=None if visible else subprocess.DEVNULL,
creationflags=creationflags,
)
async def _run_queue(host: str, port: int):
import danmu_queue
await danmu_queue.main(host=host, port=port)
async def _run_music(port: int):
import music_monitor
await music_monitor.run_monitor(port)
def _run_tts_monitor():
import tts_monitor
sys.argv = [
sys.argv[0],
"--state-file",
str(DATA_DIR / "tts_state.json"),
]
tts_monitor.main()
def _stop_children(children: list[subprocess.Popen], timeout: float = 5.0):
for child in children:
if child.poll() is None:
child.terminate()
for child in children:
if child.poll() is not None:
continue
try:
child.wait(timeout=timeout)
except subprocess.TimeoutExpired:
child.kill()
child.wait(timeout=timeout)
async def _run_all(host: str, port: int):
children: list[subprocess.Popen] = []
job = WindowsJob()
queue_task: asyncio.Task | None = None
try:
# Refuse stale/conflicting listeners before creating any helper process.
_assert_port_available(host, port)
queue_task = asyncio.create_task(_run_queue(host, port), name="queue-main")
children.append(_spawn_role("music", port, host, visible=False))
children.append(_spawn_role("tts", port, host, visible=True))
for child in children:
job.assign(child)
await queue_task
except Exception:
if queue_task is not None and not queue_task.done():
queue_task.cancel()
await asyncio.gather(queue_task, return_exceptions=True)
raise
finally:
_stop_children(children)
job.close()
def main():
ensure_runtime_dirs()
parser = argparse.ArgumentParser(description="BetterGI 直播联动统一入口")
parser.add_argument("--role", choices=["all", "queue", "music", "tts"], default="all")
parser.add_argument("--port", type=int, default=8086)
parser.add_argument("--host", default="0.0.0.0", help="Web service bind address")
args = parser.parse_args()
instance_lock = None
if args.role in {"all", "queue"}:
instance_lock = SingleInstanceLock()
if not instance_lock.acquire():
print("直播系统已经在运行,本次重复启动已拒绝。")
return 2
try:
if args.role == "tts":
_run_tts_monitor()
return 0
if args.role == "music":
asyncio.run(_run_music(args.port))
return 0
if args.role == "queue":
asyncio.run(_run_queue(args.host, args.port))
return 0
asyncio.run(_run_all(args.host, args.port))
return 0
except OSError as exc:
if getattr(exc, "winerror", None) == 10048 or getattr(exc, "errno", None) in {48, 98, 10048}:
print(f"直播端口 {args.port} 已被占用,服务未启动,也未创建辅助进程。")
return 3
raise
finally:
if instance_lock is not None:
instance_lock.close()
if __name__ == "__main__":
raise SystemExit(main())
+466
View File
@@ -0,0 +1,466 @@
"""Independent mpv audio player controlled through Windows JSON IPC."""
from __future__ import annotations
import asyncio
import json
import logging
import os
import subprocess
import threading
import time
from pathlib import Path
from typing import Any
class MpvPlayer:
def __init__(
self,
exe_path: str | Path,
logger: logging.Logger,
*,
pipe_name: str = "",
log_path: str | Path | None = None,
):
self.exe_path = Path(exe_path)
self.logger = logger
self.log_path = Path(log_path) if log_path else None
pipe_name = pipe_name or f"live_streaming_mpv_{os.getpid()}"
self.pipe_path = rf"\\.\pipe\{pipe_name}"
self.process: subprocess.Popen | None = None
self.current_url = ""
self.current_metadata: dict[str, Any] = {}
self.desired_state = "stopped"
self.generation = 0
self.started_at = 0.0
self.last_progress = 0.0
self.last_progress_at = 0.0
self.last_snapshot_at = 0.0
self.recovery_count = 0
self._request_id = 0
self._ipc_lock = asyncio.Lock()
self._pipe_state_lock = threading.Lock()
self._pipe = None
def available(self) -> bool:
return self.exe_path.is_file()
def running(self) -> bool:
return self.process is not None and self.process.poll() is None
def update_exe_path(self, exe_path: str | Path) -> None:
next_path = Path(exe_path)
if next_path == self.exe_path:
return
if self.running():
self.logger.warning(f"[mpv] 播放器路径已修改,将在进程下次重启后生效: {next_path}")
self.exe_path = next_path
async def ensure_started(self) -> bool:
if self.running():
return True
await asyncio.to_thread(self._reset_pipe_sync)
if not self.available():
self.logger.error(f"[mpv] 播放器不存在: {self.exe_path}")
return False
creationflags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0
args = [
str(self.exe_path),
"--idle=yes",
"--no-video",
"--force-window=no",
"--no-terminal",
"--msg-level=all=warn",
f"--input-ipc-server={self.pipe_path}",
"--keep-open=no",
"--audio-buffer=5",
"--cache=yes",
"--cache-secs=20",
"--demuxer-max-bytes=50MiB",
"--network-timeout=10",
]
if self.log_path:
self.log_path.parent.mkdir(parents=True, exist_ok=True)
args.append(f"--log-file={self.log_path}")
try:
self.process = subprocess.Popen(
args,
cwd=str(self.exe_path.parent),
stdin=subprocess.DEVNULL,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
creationflags=creationflags,
)
for _ in range(30):
if not self.running():
break
if await self._command(["get_property", "idle-active"], retry=False) is not None:
self.logger.info(f"[mpv] 播放服务已启动: {self.exe_path}")
return True
await asyncio.sleep(0.1)
except Exception as exc:
self.logger.error(f"[mpv] 启动失败: {exc}")
return False
def _pipe_request_sync(self, command: list[Any], request_id: int) -> Any:
request = json.dumps(
{"command": command, "request_id": request_id},
ensure_ascii=False,
).encode("utf-8") + b"\n"
pipe = None
try:
with self._pipe_state_lock:
pipe = self._pipe
if pipe is None or pipe.closed:
pipe = open(self.pipe_path, "r+b", buffering=0)
self._pipe = pipe
pipe.write(request)
deadline = time.time() + 2.5
while time.time() < deadline:
response = pipe.readline()
if not response:
continue
payload = json.loads(response.decode("utf-8", errors="replace"))
if payload.get("request_id") != request_id:
continue
if payload.get("error") != "success":
return None
if "data" not in payload:
return True
data = payload["data"]
if data is None and command and command[0] != "get_property":
return True
return data
except Exception:
self._reset_pipe_sync(pipe)
raise
return None
def _reset_pipe_sync(self, expected_pipe=None) -> None:
with self._pipe_state_lock:
pipe = self._pipe
if expected_pipe is not None and pipe is not expected_pipe:
return
self._pipe = None
if pipe is not None:
try:
pipe.close()
except Exception:
pass
async def _command(self, command: list[Any], *, retry: bool = True) -> Any:
async with self._ipc_lock:
attempts = 2 if retry else 1
for attempt in range(attempts):
self._request_id += 1
request_id = self._request_id
try:
return await asyncio.wait_for(
asyncio.to_thread(self._pipe_request_sync, command, request_id),
timeout=3.0,
)
except Exception:
self._reset_pipe_sync()
if attempt + 1 < attempts:
await asyncio.sleep(0.15)
return None
def _clear_current(self) -> None:
self.current_url = ""
self.current_metadata = {}
self.started_at = 0.0
self.last_progress = 0.0
self.last_progress_at = 0.0
self.recovery_count = 0
@staticmethod
def _path_matches(expected: str, actual: str) -> bool:
expected = str(expected or "").strip()
actual = str(actual or "").strip()
if not expected or not actual:
return False
if expected == actual:
return True
if expected.lower().startswith(("http://", "https://")):
return False
try:
actual_path = actual[8:] if actual.lower().startswith("file:///") else actual
return Path(expected).resolve() == Path(actual_path).resolve()
except Exception:
return False
async def _wait_until_loaded(
self,
*,
expected_generation: int,
start_at: float,
timeout: float = 12.0,
) -> bool:
deadline = time.time() + max(2.0, timeout)
expected_url = self.current_url
while time.time() < deadline:
if expected_generation != self.generation or self.desired_state != "playing":
return False
if not self.running():
return False
path = await self._command(["get_property", "path"], retry=False)
idle = await self._command(["get_property", "idle-active"], retry=False)
if self._path_matches(expected_url, str(path or "")) and idle is False:
if start_at > 0:
seek_result = None
seek_deadline = min(deadline, time.time() + 4.0)
while time.time() < seek_deadline:
duration = await self._command(["get_property", "duration"], retry=False)
try:
duration_value = max(0.0, float(duration or 0.0))
except (TypeError, ValueError):
duration_value = 0.0
if duration_value > 0:
seek_result = await self._command(["seek", float(start_at), "absolute+exact"])
if seek_result is not None:
break
await asyncio.sleep(0.2)
if seek_result is None:
self.logger.warning(f"[mpv] 续播定位失败: {start_at:.1f}")
return False
if await self._command(["set_property", "pause", False]) is None:
return False
progress_deadline = min(deadline, time.time() + 4.0)
baseline = max(0.0, float(start_at or 0.0))
while time.time() < progress_deadline:
progress = await self._command(["get_property", "time-pos"], retry=False)
duration = await self._command(["get_property", "duration"], retry=False)
try:
progress_value = max(0.0, float(progress or 0.0))
except (TypeError, ValueError):
progress_value = 0.0
try:
duration_value = max(0.0, float(duration or 0.0))
except (TypeError, ValueError):
duration_value = 0.0
idle_now = await self._command(["get_property", "idle-active"], retry=False)
if idle_now is False and (progress_value > 0.05 or duration_value > 0 or baseline > 0):
now = time.time()
self.started_at = now
self.last_progress = max(baseline, progress_value)
self.last_progress_at = now
return True
await asyncio.sleep(0.15)
await asyncio.sleep(0.15)
self.logger.warning(
f"[mpv] 音频加载超时,未进入可播放状态: "
f"{self.current_metadata.get('name') or self.current_url}"
)
return False
async def _load_current(self, *, start_at: float, expected_generation: int) -> bool:
if expected_generation != self.generation or self.desired_state != "playing" or not self.current_url:
return False
if not await self.ensure_started():
return False
if expected_generation != self.generation or self.desired_state != "playing":
return False
result = await self._command(["loadfile", self.current_url, "replace"])
if result is None or expected_generation != self.generation or self.desired_state != "playing":
return False
return await self._wait_until_loaded(
expected_generation=expected_generation,
start_at=start_at,
)
async def play(self, url: str, metadata: dict[str, Any], *, start_at: float = 0.0) -> bool:
self.generation += 1
generation = self.generation
self.desired_state = "playing"
self.current_url = str(url)
self.current_metadata = dict(metadata)
self.recovery_count = 0
ok = await self._load_current(start_at=start_at, expected_generation=generation)
if not ok and generation == self.generation:
self.desired_state = "stopped"
self._clear_current()
return ok
async def pause(self) -> bool:
if not self.current_url:
return False
self.desired_state = "paused"
if not self.running():
return True
return await self._command(["set_property", "pause", True]) is not None
async def resume(self) -> bool:
if not self.current_url:
return False
self.desired_state = "playing"
self.last_progress_at = time.time()
if not self.running():
return await self._load_current(start_at=self.last_progress, expected_generation=self.generation)
return await self._command(["set_property", "pause", False]) is not None
async def stop(self) -> bool:
self.generation += 1
self.desired_state = "stopped"
self._clear_current()
if not self.running():
return True
return await self._command(["stop"]) is not None
async def close(self) -> None:
self.generation += 1
self.desired_state = "stopped"
self._clear_current()
if self.running():
await self._command(["quit"])
await asyncio.sleep(0.15)
if self.running():
self.process.terminate()
try:
await asyncio.to_thread(self.process.wait, 2)
except Exception:
if self.running():
self.process.kill()
self.process = None
await asyncio.to_thread(self._reset_pipe_sync)
async def snapshot(self) -> dict[str, Any]:
now = time.time()
if not self.running():
return self._snapshot_payload(False, 0.0, 0.0, True, False, False, "", now)
progress, duration, paused, idle, eof, path = await asyncio.gather(
self._command(["get_property", "time-pos"]),
self._command(["get_property", "duration"]),
self._command(["get_property", "pause"]),
self._command(["get_property", "idle-active"]),
self._command(["get_property", "eof-reached"]),
self._command(["get_property", "path"]),
)
try:
progress_value = max(0.0, float(progress or 0.0))
except (TypeError, ValueError):
progress_value = 0.0
try:
duration_value = max(0.0, float(duration or self.current_metadata.get("duration_sec", 0) or 0))
except (TypeError, ValueError):
duration_value = 0.0
idle_value = bool(idle) if idle is not None else not bool(self.current_url)
paused_value = bool(paused)
eof_value = bool(eof)
playing = (
self.desired_state == "playing"
and bool(self.current_url)
and not paused_value
and not idle_value
and not eof_value
)
if progress_value > self.last_progress + 0.2:
self.last_progress = progress_value
self.last_progress_at = now
self.recovery_count = 0
self.last_snapshot_at = now
return self._snapshot_payload(
playing,
progress_value,
duration_value,
idle_value,
paused_value,
eof_value,
str(path or ""),
now,
)
def _snapshot_payload(
self,
playing: bool,
progress: float,
duration: float,
idle: bool,
paused: bool,
eof: bool,
path: str,
now: float,
) -> dict[str, Any]:
metadata = self.current_metadata
return {
"playing": playing,
"paused": paused,
"idle": idle,
"eof": eof,
"desired_state": self.desired_state,
"generation": self.generation,
"path": path,
"current": {
"id": str(metadata.get("id", "")),
"title": metadata.get("name") or ("暂无歌曲" if idle else "正在加载"),
"artist": metadata.get("artist") or "mpv",
"cover": metadata.get("cover", ""),
"cover_hash": metadata.get("cover_hash", ""),
"duration": duration,
"progress": progress,
"source": "mpv",
},
"playlist": [],
"requests": [],
"monitor": {
"online": self.running(),
"source": "mpv.ipc",
"platform": "mpv",
"updated_at": now,
"targets": [],
"allow_all": False,
},
}
def mark_ended(self) -> None:
self.generation += 1
self.desired_state = "stopped"
self._clear_current()
async def maintain(self, *, stall_seconds: float = 12.0) -> dict[str, Any]:
"""Recover only unexpected failures while the desired state is playing."""
if self.desired_state != "playing" or not self.current_url:
return {"action": "none"}
generation = self.generation
now = time.time()
if not self.running():
progress = self.last_progress
self.process = None
if await self._load_current(start_at=progress, expected_generation=generation):
self.logger.warning(f"[mpv] 进程退出后已从 {progress:.1f} 秒恢复")
return {"action": "process_restarted", "progress": progress}
return {"action": "failed", "reason": "process_restart_failed"}
state = await self.snapshot()
if generation != self.generation or self.desired_state != "playing":
return {"action": "superseded", "snapshot": state}
progress = float((state.get("current") or {}).get("progress", 0) or 0)
duration = float((state.get("current") or {}).get("duration", 0) or 0)
effective_progress = max(progress, self.last_progress)
near_end = duration > 0 and effective_progress >= max(0.0, duration - 2.0)
unloaded_after_progress = (
state.get("idle")
and not str(state.get("path") or "")
and effective_progress >= 0.5
)
if state.get("eof") or (state.get("idle") and near_end) or unloaded_after_progress:
self.mark_ended()
return {"action": "ended", "progress": effective_progress, "duration": duration, "snapshot": state}
loading_grace = now - self.started_at < 2.5
if state.get("paused") and not loading_grace:
if await self._command(["set_property", "pause", False]) is not None:
self.last_progress_at = now
self.logger.warning("[mpv] 检测到非预期暂停,已自动继续播放")
return {"action": "resumed", "progress": progress, "snapshot": state}
if state.get("idle") and not loading_grace:
return {"action": "reload_required", "reason": "unexpected_idle", "progress": progress, "snapshot": state}
if self.last_progress_at and not loading_grace and now - self.last_progress_at >= max(3.0, stall_seconds):
self.recovery_count += 1
if self.recovery_count == 1:
await self._command(["set_property", "pause", False])
self.last_progress_at = now
return {"action": "unstalled", "progress": progress, "snapshot": state}
return {"action": "reload_required", "reason": "stalled", "progress": progress, "snapshot": state}
return {"action": "none", "snapshot": state}
+965
View File
@@ -0,0 +1,965 @@
"""
音乐播放器 -> 直播间 UI 音乐状态同步
通过 Windows 10/11 的 SMTCSystem Media Transport Controls)读取当前媒体会话,
并推送到 BGI 直播间 Web 服务。默认偏向网易云音乐,但可在 config/config.json
的 music_monitor 节点或命令行中配置目标播放器。
播放状态判断(参考 now-playing-service/NeteaseMusicService.cs):
- 优先用 Windows 音频会话的峰值音量(volume>0 = Playing),最贴近"是否真在出声"
- volume=0 但 UIA 进度最近 1.5s 内变化 → 仍视为 Playing(静音播放/拖进度条场景)
- 否则 Paused
封面防闪:
- 只在 SMTC 缩略图 hash 变化时才重写 music_cover.jpg(原子 tmp→replace
- 同一 hash 不重写文件,前端不会因文件变化触发 reload
- 标题/歌手来自 cloudmusic 窗口标题(已稳定,不改)
"""
import argparse
import asyncio
import hashlib
import ctypes
import json
import os
import re
import subprocess
import sys
import time
import traceback
from pathlib import Path
from typing import Any
APP_DIR = Path(__file__).resolve().parent
if str(APP_DIR) not in sys.path:
sys.path.insert(0, str(APP_DIR))
import aiohttp
import winrt.windows.media.control as wmc
import winrt.windows.storage.streams as streams
try:
import uiautomation as uia
except Exception:
uia = None
# pycaw: 读取 Windows 音频会话峰值音量(参考 now-playing-service 的 CSCore.AudioMeterInformation
try:
from pycaw.pycaw import AudioUtilities # type: ignore
_PYCAW_OK = True
except Exception:
_PYCAW_OK = False
from core.runtime_paths import CONFIG_DIR, DATA_DIR, WEB_DIR, ensure_runtime_dirs
ensure_runtime_dirs()
CONFIG_FILE = CONFIG_DIR / "config.json"
MUSIC_FILE = DATA_DIR / "music_state.json"
COVER_FILE = WEB_DIR / "music_cover.jpg"
DEFAULT_TARGETS = [
"网易云音乐",
"Netease",
"CloudMusic",
"cloudmusic",
"YesPlayMusic",
"Listen1",
"QQMusic",
"qqmusic",
"spotify",
]
DEFAULT_MONITOR_CONFIG = {
"platform": "netease",
"targets": ["网易云音乐", "Netease", "CloudMusic", "cloudmusic"],
"allow_all": False,
"interval_sec": 1.0,
"holdover_ms": 1500,
"prefer_playing": True,
"keep_last_when_none": True,
"cover_enabled": True,
"auto_resume_enabled": False,
"auto_resume_interval_sec": 3,
"auto_resume_stall_sec": 10,
"extra_filter": "",
}
def load_monitor_config() -> dict[str, Any]:
cfg = dict(DEFAULT_MONITOR_CONFIG)
if CONFIG_FILE.exists():
try:
data = json.loads(CONFIG_FILE.read_text(encoding="utf-8"))
if isinstance(data.get("music_monitor"), dict):
cfg.update(data["music_monitor"])
except Exception:
pass
targets = cfg.get("targets") or []
if isinstance(targets, str):
targets = [x.strip() for x in targets.replace("", ",").split(",") if x.strip()]
extra = str(cfg.get("extra_filter", "")).strip()
if extra:
targets.append(extra)
cfg["targets"] = [str(x).strip() for x in targets if str(x).strip()]
return cfg
def session_identity(session) -> str:
display = getattr(session, "source_app_display_name", "") or ""
aumid = getattr(session, "source_app_user_model_id", "") or ""
return f"{display} {aumid}".strip()
def is_target_app(session, targets: list[str], allow_all: bool) -> bool:
if allow_all or not targets:
return True
combined = session_identity(session).lower()
return any(target.lower() in combined for target in targets)
def playback_status_int(session) -> int:
try:
return int(session.get_playback_info().playback_status)
except Exception:
return 0
def is_playing(session) -> bool:
return playback_status_int(session) == 4
async def choose_session(mgr, cfg: dict[str, Any]):
"""从全部媒体会话中选择最合适的目标会话。"""
targets = cfg.get("targets", [])
allow_all = bool(cfg.get("allow_all", False))
prefer_playing = bool(cfg.get("prefer_playing", True))
sessions = list(mgr.get_sessions())
candidates = [s for s in sessions if is_target_app(s, targets, allow_all)]
if not candidates:
current = mgr.get_current_session()
if current and is_target_app(current, targets, allow_all):
candidates = [current]
if not candidates:
return None
# 当前系统媒体会话通常最能代表用户正在操作的播放器;优先选择它,
# 避免网易云残留多个 SMTC 会话时反复读到已停止更新的旧会话。
current = mgr.get_current_session()
if current and current in candidates and (not prefer_playing or is_playing(current)):
return current
if prefer_playing:
playing = [s for s in candidates if is_playing(s)]
if playing:
return playing[0]
if current and current in candidates:
return current
return candidates[0]
async def read_thumbnail(thumb_ref):
if thumb_ref is None:
return None
try:
stream = await thumb_ref.open_read_async()
size = stream.size
if size <= 0 or size > 5 * 1024 * 1024:
return None
buffer = streams.Buffer(size)
await stream.read_async(buffer, size, streams.InputStreamOptions.READ_AHEAD)
data = bytes(buffer)
return data if len(data) > 100 else None
except Exception:
return None
def write_cover_atomic(data: bytes, previous_hash: str = "") -> tuple[str, str]:
cover_hash = hashlib.md5(data).hexdigest()
if cover_hash == previous_hash and COVER_FILE.exists():
return "/music_cover.jpg", cover_hash
COVER_FILE.parent.mkdir(parents=True, exist_ok=True)
tmp = COVER_FILE.with_suffix(".tmp")
tmp.write_bytes(data)
tmp.replace(COVER_FILE)
return "/music_cover.jpg", cover_hash
def parse_time_to_seconds(text: str) -> int | None:
parts = str(text).strip().split(":")
if len(parts) not in (2, 3):
return None
try:
nums = [int(x) for x in parts]
except ValueError:
return None
if any(x < 0 for x in nums) or nums[-1] >= 60:
return None
if len(nums) == 2:
return nums[0] * 60 + nums[1]
if nums[1] >= 60:
return None
return nums[0] * 3600 + nums[1] * 60 + nums[2]
def parse_progress_text(text: str) -> tuple[int, int] | None:
cleaned = str(text).replace(" ", "")
m = re.search(r"(\d{1,2}:\d{2}(?::\d{2})?)\s*[/|/|]\s*(\d{1,2}:\d{2}(?::\d{2})?)", cleaned)
if not m:
return None
current = parse_time_to_seconds(m.group(1))
total = parse_time_to_seconds(m.group(2))
if current is None or total is None or total <= 0 or current > total + 2:
return None
return max(0, current), max(0, total)
def _cloudmusic_pids() -> list[int]:
if not hasattr(ctypes, "windll"):
return []
try:
output = subprocess.check_output(
["tasklist", "/FI", "IMAGENAME eq cloudmusic.exe", "/FO", "CSV", "/NH"],
text=True, encoding="gbk", errors="ignore", creationflags=0x08000000,
)
pids = []
for line in output.splitlines():
parts = [p.strip().strip('"') for p in line.split(",")]
if len(parts) >= 2 and parts[0].lower() == "cloudmusic.exe":
try:
pids.append(int(parts[1]))
except ValueError:
pass
return pids
except Exception:
return []
def _window_titles_by_pids(pids: list[int]) -> list[str]:
if not pids or not hasattr(ctypes, "windll"):
return []
titles = []
user32 = ctypes.windll.user32
EnumWindowsProc = ctypes.WINFUNCTYPE(ctypes.c_bool, ctypes.c_void_p, ctypes.c_void_p)
def enum_proc(hwnd, lparam):
try:
if not user32.IsWindowVisible(hwnd):
return True
pid = ctypes.c_ulong()
user32.GetWindowThreadProcessId(hwnd, ctypes.byref(pid))
if pid.value not in pids:
return True
length = user32.GetWindowTextLengthW(hwnd)
if length <= 0:
return True
buf = ctypes.create_unicode_buffer(length + 1)
user32.GetWindowTextW(hwnd, buf, length + 1)
title = (buf.value or "").strip()
if title:
titles.append(title)
except Exception:
pass
return True
try:
user32.EnumWindows(EnumWindowsProc(enum_proc), 0)
except Exception:
pass
return titles
def read_netease_window_title() -> str:
"""参考 now-playing-service:从 cloudmusic 进程窗口标题取“歌名 - 歌手”。"""
for title in _window_titles_by_pids(_cloudmusic_pids()):
if " - " in title and "MediaPlayer" not in title:
return title.replace("/", " / ").strip()
return ""
def read_netease_progress_uia() -> tuple[int, int] | None:
"""参考 now-playing-service:只在 cloudmusic 播放窗口子树中解析 MM:SS / MM:SS。"""
if uia is None:
return None
try:
pids = set(_cloudmusic_pids())
root = uia.GetRootControl()
for win in root.GetChildren():
try:
pid = int(getattr(win, "ProcessId", 0) or 0)
name = (getattr(win, "Name", "") or "")
if pids and pid not in pids:
continue
if " - " not in name and "cloudmusic" not in (getattr(win, "ClassName", "") or "").lower():
continue
stack = list(win.GetChildren())
deadline = time.time() + 0.3
while stack and time.time() < deadline:
ctrl = stack.pop(0)
text_value = (getattr(ctrl, "Name", "") or "").strip()
parsed = parse_progress_text(text_value)
if parsed:
return parsed
try:
stack.extend(ctrl.GetChildren())
except Exception:
pass
except Exception:
continue
except Exception:
return None
return None
_LAST_NETEASE_TITLE = ""
_LAST_NETEASE_AT = 0.0
_LAST_PROGRESS_SECONDS = -1
_LAST_PROGRESS_CHANGE_AT = 0.0
_LAST_VOLUME_PEAK = 0.0
_LAST_PLAYING = False
_LAST_COVER_HASH = ""
_LAST_COVER_SOURCE = ""
_COVER_CACHE: dict[str, tuple[str, str]] = {}
_COVER_TASKS: dict[str, asyncio.Task] = {}
_SMTC_DISABLED_UNTIL = 0.0
_LAST_AUTO_RESUME_AT = 0.0
_AUTO_RESUME_TRACK_KEY = ""
_AUTO_RESUME_LAST_PROGRESS = -1.0
_AUTO_RESUME_LAST_PROGRESS_AT = 0.0
_AUTO_RESUME_PAUSED_SINCE = 0.0
_AUTO_RESUME_FAILURES = 0
async def get_target_smtc_session(cfg: dict[str, Any]):
mgr = await wmc.GlobalSystemMediaTransportControlsSessionManager.request_async()
return await choose_session(mgr, cfg)
async def try_resume_playback(cfg: dict[str, Any], force: bool = False) -> bool:
"""通过 SMTC 尝试恢复播放;暂停/停滞时允许重复发送播放命令。"""
global _SMTC_DISABLED_UNTIL
if time.time() < _SMTC_DISABLED_UNTIL:
return False
async def _inner() -> bool:
session = await get_target_smtc_session(cfg)
if session is None:
return False
if is_playing(session) and not force:
return False
result = await session.try_play_async()
return bool(result)
def _run():
return asyncio.run(_inner())
try:
return await asyncio.wait_for(asyncio.to_thread(_run), timeout=2.5)
except Exception:
_SMTC_DISABLED_UNTIL = time.time() + 5
return False
async def maybe_auto_resume(info: dict[str, Any] | None, cfg: dict[str, Any]) -> str:
"""定期恢复明确暂停,也检测“状态为播放但进度长时间不动”的假播放。"""
global _LAST_AUTO_RESUME_AT, _AUTO_RESUME_TRACK_KEY
global _AUTO_RESUME_LAST_PROGRESS, _AUTO_RESUME_LAST_PROGRESS_AT
global _AUTO_RESUME_PAUSED_SINCE, _AUTO_RESUME_FAILURES
if not bool(cfg.get("auto_resume_enabled", False)):
return ""
if info is None or not info.get("title"):
_AUTO_RESUME_TRACK_KEY = ""
_AUTO_RESUME_LAST_PROGRESS = -1.0
_AUTO_RESUME_LAST_PROGRESS_AT = 0.0
_AUTO_RESUME_PAUSED_SINCE = 0.0
_AUTO_RESUME_FAILURES = 0
return ""
now = time.time()
title = str(info.get("title") or "").strip()
artist = str(info.get("artist") or "").strip()
track_key = f"{title}\n{artist}"
progress = max(0.0, float(info.get("progress") or 0))
duration = max(0.0, float(info.get("duration") or 0))
playing = bool(info.get("playing"))
if track_key != _AUTO_RESUME_TRACK_KEY:
_AUTO_RESUME_TRACK_KEY = track_key
_AUTO_RESUME_LAST_PROGRESS = progress
_AUTO_RESUME_LAST_PROGRESS_AT = now
_AUTO_RESUME_PAUSED_SINCE = 0.0 if playing else now
_AUTO_RESUME_FAILURES = 0
return ""
if playing:
_AUTO_RESUME_PAUSED_SINCE = 0.0
elif not _AUTO_RESUME_PAUSED_SINCE:
_AUTO_RESUME_PAUSED_SINCE = now
if abs(progress - _AUTO_RESUME_LAST_PROGRESS) >= 0.5:
_AUTO_RESUME_LAST_PROGRESS = progress
_AUTO_RESUME_LAST_PROGRESS_AT = now
_AUTO_RESUME_FAILURES = 0
return ""
interval_sec = max(1.0, float(cfg.get("auto_resume_interval_sec", 3) or 3))
stall_sec = max(interval_sec, float(cfg.get("auto_resume_stall_sec", 10) or 10))
stalled = (
playing
and progress > 0
and (duration <= 0 or progress < max(0, duration - 2))
and now - _AUTO_RESUME_LAST_PROGRESS_AT >= stall_sec
)
if playing and not stalled:
return ""
if now - _LAST_AUTO_RESUME_AT < interval_sec:
return ""
_LAST_AUTO_RESUME_AT = now
# 明确暂停也强制发送 try_play,避免 SMTC 状态缓存或会话切换导致第一次命令被吞掉。
resumed = await try_resume_playback(cfg, force=True)
if resumed:
_AUTO_RESUME_LAST_PROGRESS_AT = now
_AUTO_RESUME_FAILURES = 0
return "stalled" if stalled else "paused"
_AUTO_RESUME_FAILURES += 1
# 连续恢复失败时缩短 SMTC 禁用窗口,下一轮重新获取会话并继续尝试。
if _AUTO_RESUME_FAILURES >= 2:
global _SMTC_DISABLED_UNTIL
_SMTC_DISABLED_UNTIL = min(_SMTC_DISABLED_UNTIL, now + 1)
return ""
async def read_smtc_snapshot(cfg: dict[str, Any]) -> dict[str, Any] | None:
"""在隔离线程里读取 SMTC,避免 WinRT 偶发卡死拖住 MusicMonitor 主循环。"""
global _SMTC_DISABLED_UNTIL
if time.time() < _SMTC_DISABLED_UNTIL:
return None
async def _inner():
session = await get_target_smtc_session(cfg)
if session is None:
return None
props = await session.try_get_media_properties_async()
playback = session.get_playback_info()
timeline = session.get_timeline_properties()
thumb_data = None
if bool(cfg.get("cover_enabled", True)):
thumb_data = await read_thumbnail(props.thumbnail)
return {
"title": props.title or "",
"artist": props.artist or "",
"duration": max(0, timeline.end_time.total_seconds()),
"progress": max(0, timeline.position.total_seconds()),
"playing": int(playback.playback_status) == 4,
"source": session_identity(session) or "smtc",
"thumb_data": thumb_data,
}
def _run():
return asyncio.run(_inner())
try:
return await asyncio.wait_for(asyncio.to_thread(_run), timeout=2.5)
except Exception:
_SMTC_DISABLED_UNTIL = time.time() + 5
return None
def _get_cloudmusic_audio_peak() -> float:
"""参考 now-playing-service: 累加 cloudmusic 所有音频会话的峰值音量。
volume>0 = 真在出声 = Playing。pycaw 不可用时返回 -1 表示未知。"""
if not _PYCAW_OK:
return -1.0
try:
from pycaw.pycaw import IAudioMeterInformation # type: ignore
total = 0.0
sessions = AudioUtilities.GetAllSessions()
for sess in sessions:
try:
proc = getattr(sess, "Process", None)
if proc is None:
continue
# pycaw 的 Process.name 是方法不是属性, 要调用
name_attr = getattr(proc, "name", None)
if callable(name_attr):
pname = name_attr()
else:
pname = str(name_attr or "")
pname = (pname or "").lower()
if "cloudmusic" in pname:
meter = sess._ctl.QueryInterface(IAudioMeterInformation)
total += meter.GetPeakValue()
except Exception:
continue
return total
except Exception:
return -1.0
def _decide_playing(volume_peak: float, progress_changed_recently: bool, holdover_sec: float) -> bool:
"""参考 now-playing-service: volume>0 → Playing; volume=0 但进度最近变化 → Playing; 否则保持/暂停。"""
global _LAST_PLAYING
if volume_peak > 0.00001:
_LAST_PLAYING = True
return True
if volume_peak < 0:
# pycaw 不可用,回退到进度判断
if progress_changed_recently:
_LAST_PLAYING = True
return True
_LAST_PLAYING = False
return False
# volume=0
if progress_changed_recently:
# 进度在动但没声音 → 静音播放,仍算 Playing
_LAST_PLAYING = True
return True
_LAST_PLAYING = False
return False
def _split_title_artist(window_title: str) -> tuple[str, str]:
title = (window_title or "").strip()
if " - " in title:
song, artist = title.split(" - ", 1)
return song.strip() or title, artist.strip()
return title, ""
import websockets as _ws_mod
def fiber_store_extract_js() -> str:
return r'''
function _ensureStore() {
try {
if (window._reduxStore) return true;
const rootEl = document.querySelector('#root');
const root = window._fiberRoot || (rootEl && rootEl._reactRootContainer && rootEl._reactRootContainer._internalRoot);
if (!root) return false;
let queue = [root.current || root];
let visited = 0;
while (queue.length > 0) {
let node = queue.shift();
if (!node) continue;
visited++;
if (visited > 20000) break;
if (node.memoizedProps && node.memoizedProps.store) { window._reduxStore = node.memoizedProps.store; return true; }
if (node.stateNode && node.stateNode.store) { window._reduxStore = node.stateNode.store; return true; }
let child = node.child;
while (child) { queue.push(child); child = child.sibling; }
}
return false;
} catch(err) { return false; }
}
'''
async def get_cdp_ws_url(port: int = 9222) -> str:
try:
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=1.2)) as session:
async with session.get(f"http://127.0.0.1:{port}/json") as resp:
if resp.status != 200:
return ""
targets = await resp.json(content_type=None)
for target in targets:
text = (str(target.get("url", "")) + " " + str(target.get("title", ""))).lower()
if target.get("type") == "page" and target.get("webSocketDebuggerUrl") and ("orpheus" in text or "music.163.com" in text):
return target.get("webSocketDebuggerUrl", "")
for target in targets:
if target.get("type") == "page" and target.get("webSocketDebuggerUrl"):
return target.get("webSocketDebuggerUrl", "")
except Exception:
return ""
return ""
async def read_netease_cdp_state(cfg: dict[str, Any], previous_hash: str = "") -> dict[str, Any] | None:
"""通过 CDP 读网易云 Redux 状态。只取 title/artist/progress/duration/picUrl
不取 playingplaying 交给音量峰值判断)。"""
port = int(((cfg.get("request_player") or {}).get("cdp_port", 9222)) or 9222)
ws_url = await get_cdp_ws_url(port)
if not ws_url:
return None
script = fiber_store_extract_js() + r'''
(function(){
if(!_ensureStore()) return null;
const state = window._reduxStore.getState();
const playing = state.playing || {};
const list = (state.playingList && state.playingList.curPlayingList) || [];
const id = playing.resourceTrackId || playing.onlineResourceId || playing.resourceId || playing.trackId;
function normId(x){ return x == null ? '' : String(x); }
let item = null;
if(id) item = list.find(x => normId(x.id || x.trackId || x.resourceId) === normId(id));
if(!item && list.length === 1) item = list[0];
const track = (item && (item.track || item.resource || item)) || {};
const artists = track.artists || track.ar || item?.artists || item?.ar || [];
let artist = '';
if(Array.isArray(artists)) artist = artists.map(a => a && a.name ? a.name : '').filter(Boolean).join('/');
else if(typeof artists === 'string') artist = artists;
const album = track.album || track.al || item?.album || item?.al || {};
const picUrl = album.picUrl || album.blurPicUrl || track.picUrl || item?.picUrl || '';
const durationMs = Number(track.duration || track.dt || item?.duration || item?.dt || 0);
const positionMs = Number(playing.position || playing.currentTime || playing.progress || 0);
return {
id: normId(id || item?.id || track.id),
title: track.name || item?.name || '',
artist: artist || '',
picUrl: picUrl || '',
duration: durationMs > 10000 ? durationMs / 1000 : durationMs,
progress: positionMs > 10000 ? positionMs / 1000 : positionMs,
};
})()
'''
try:
async with _ws_mod.connect(ws_url, open_timeout=1.5, close_timeout=0.5) as ws:
await ws.send(json.dumps({
"id": 1,
"method": "Runtime.evaluate",
"params": {"expression": script, "returnByValue": True, "awaitPromise": True},
}))
deadline = time.time() + 2
while time.time() < deadline:
raw = await asyncio.wait_for(ws.recv(), timeout=max(0.1, deadline - time.time()))
msg = json.loads(raw)
if msg.get("id") != 1:
continue
value = (((msg.get("result") or {}).get("result") or {}).get("value"))
if isinstance(value, dict) and (value.get("title") or value.get("id")):
return value
return None
except Exception:
return None
return None
async def download_cover_url(pic_url: str, previous_hash: str = "") -> tuple[str, str]:
global _LAST_COVER_SOURCE
if not pic_url:
return "", previous_hash
if pic_url.startswith("http://"):
pic_url = "https://" + pic_url[7:]
if pic_url == _LAST_COVER_SOURCE and previous_hash and COVER_FILE.exists():
return "/music_cover.jpg", previous_hash
try:
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=3)) as session:
async with session.get(pic_url) as resp:
if resp.status != 200:
return "", previous_hash
data = await resp.read()
if data:
cover_url, cover_hash = write_cover_atomic(data, previous_hash)
_LAST_COVER_SOURCE = pic_url
return cover_url, cover_hash
except Exception:
return "", previous_hash
return "", previous_hash
def cover_cache_key(title: str, artist: str) -> str:
return re.sub(r"\s+", " ", f"{title} - {artist}".strip().lower())
async def fetch_cover_task(key: str, title: str, artist: str):
"""后台补封面:失败也不能影响 MusicMonitor 主循环。"""
try:
keyword = " ".join(x for x in [title, artist] if x).strip()
if not keyword:
return
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36",
"Referer": "https://music.163.com/",
"Content-Type": "application/x-www-form-urlencoded",
"Cookie": "os=pc; appver=2.9.8;",
"X-Real-IP": "118.88.88.88",
"X-Forwarded-For": "118.88.88.88",
}
async with aiohttp.ClientSession(headers=headers, timeout=aiohttp.ClientTimeout(total=4)) as session:
async with session.post(
"https://music.163.com/api/search/get/web",
data={"s": keyword, "type": "1", "limit": "1", "offset": "0"},
) as resp:
result = await resp.json(content_type=None)
songs = ((result.get("result") or {}).get("songs") or [])
if not songs:
return
album = songs[0].get("album") or {}
pic_url = album.get("picUrl") or album.get("blurPicUrl") or ""
if not pic_url:
return
if pic_url.startswith("http://"):
pic_url = "https://" + pic_url[7:]
async with session.get(pic_url) as img_resp:
if img_resp.status != 200:
return
data = await img_resp.read()
if data:
_COVER_CACHE[key] = write_cover_atomic(data)
except Exception:
pass
finally:
_COVER_TASKS.pop(key, None)
def apply_cached_or_schedule_cover(info: dict[str, Any], cfg: dict[str, Any]):
if not bool(cfg.get("cover_enabled", True)):
return
key = cover_cache_key(info.get("title", ""), info.get("artist", ""))
if not key:
return
cached = _COVER_CACHE.get(key)
if cached:
info["cover"], info["cover_hash"] = cached
return
task = _COVER_TASKS.get(key)
if task is None or task.done():
try:
_COVER_TASKS[key] = asyncio.create_task(fetch_cover_task(key, info.get("title", ""), info.get("artist", "")))
except RuntimeError:
pass
def apply_netease_window_fallback(info: dict[str, Any], cfg: dict[str, Any]) -> dict[str, Any]:
"""用窗口标题修正歌曲名/艺人,用音量峰值+进度变化判断播放状态。
参考 now-playing-service/NeteaseMusicService.cs。"""
global _LAST_NETEASE_TITLE, _LAST_NETEASE_AT, _LAST_PROGRESS_SECONDS, _LAST_PROGRESS_CHANGE_AT
holdover = max(0, int(cfg.get("holdover_ms", 1500))) / 1000
win_title = read_netease_window_title()
now = time.time()
if win_title:
_LAST_NETEASE_TITLE = win_title
_LAST_NETEASE_AT = now
elif _LAST_NETEASE_TITLE and now - _LAST_NETEASE_AT <= holdover:
win_title = _LAST_NETEASE_TITLE
if win_title:
title, artist = _split_title_artist(win_title)
if title:
info["title"] = title
if artist:
info["artist"] = artist
info["source"] = "cloudmusic.window"
# 进度
parsed = read_netease_progress_uia()
if parsed:
progress, duration = parsed
info["progress"] = progress
info["duration"] = duration
if progress != _LAST_PROGRESS_SECONDS:
_LAST_PROGRESS_SECONDS = progress
_LAST_PROGRESS_CHANGE_AT = now
progress_changed_recently = (now - _LAST_PROGRESS_CHANGE_AT) <= holdover
# 播放状态:音量峰值优先,进度兜底
volume_peak = _get_cloudmusic_audio_peak()
info["playing"] = _decide_playing(volume_peak, progress_changed_recently, holdover)
return info
async def get_media_info(cfg: dict[str, Any], previous_cover_hash: str = ""):
"""网易云已开启 SMTC 后:优先信任 SMTC 的标题、封面、播放状态;CDP/窗口标题只做兜底。"""
is_netease = any("cloudmusic" in str(t).lower() for t in cfg.get("targets", []))
smtc = await read_smtc_snapshot(cfg)
info = {
"title": "",
"artist": "",
"duration": 0,
"progress": 0,
"playing": False,
"cover": "",
"cover_hash": previous_cover_hash,
"source": "smtc",
"updated_at": time.time(),
}
cover_url = ""
cover_hash = previous_cover_hash
if smtc is not None:
info["title"] = smtc.get("title", "")
info["artist"] = smtc.get("artist", "")
info["duration"] = float(smtc.get("duration") or 0)
info["progress"] = float(smtc.get("progress") or 0)
info["source"] = smtc.get("source") or "smtc"
info["playing"] = bool(smtc.get("playing"))
thumb_data = smtc.get("thumb_data")
if thumb_data:
cover_url, cover_hash = write_cover_atomic(thumb_data, previous_cover_hash)
if is_netease:
# SMTC 现在是主数据源;如果 SMTC 某些字段缺失,再用 CDP/窗口标题补齐。
# 现在网易云已开启 SMTC,音乐显示主链路不再碰 CDP/WebSocket,避免卡主循环。
# 如果 SMTC 暂时没给标题,再用窗口标题兜底;封面则保留上一轮,不再网络下载。
if not info.get("title"):
info = apply_netease_window_fallback(info, cfg)
if not info.get("title"):
return None
elif smtc is None:
return None
# 封面绝不因为某一轮没读到就清空,避免前台闪烁。
if cover_url:
info["cover"] = cover_url
elif previous_cover_hash and COVER_FILE.exists():
info["cover"] = "/music_cover.jpg"
else:
# SMTC/窗口标题没有封面时,按“歌名 + 歌手”异步补封面;成功后下一轮自动显示。
apply_cached_or_schedule_cover(info, cfg)
if not info.get("cover"):
info["cover"] = ""
info["cover_hash"] = info.get("cover_hash") or cover_hash
return info
def default_state() -> dict[str, Any]:
return {
"playing": False,
"current": {
"title": "暂无歌曲",
"artist": "未接入音乐源",
"cover": "",
"duration": 0,
"progress": 0,
"source": "",
},
"playlist": [],
"requests": [],
"monitor": {"online": False, "source": ""},
}
def load_state() -> dict[str, Any]:
if MUSIC_FILE.exists():
try:
state = json.loads(MUSIC_FILE.read_text(encoding="utf-8"))
if isinstance(state, dict):
return state
except Exception:
pass
return default_state()
async def save_state(state: dict[str, Any]):
MUSIC_FILE.parent.mkdir(parents=True, exist_ok=True)
tmp = MUSIC_FILE.with_suffix(".tmp")
tmp.write_text(json.dumps(state, ensure_ascii=False, indent=2), encoding="utf-8")
tmp.replace(MUSIC_FILE)
async def post_state(state: dict[str, Any], api_url: str) -> bool:
try:
async with aiohttp.ClientSession() as session:
async with session.post(api_url, json=state) as resp:
return resp.status == 200
except Exception:
return False
async def run_monitor(port: int, interval: float | None = None, cli_targets=None, allow_all: bool | None = None):
api_url = f"http://localhost:{port}/api/music"
state = load_state()
last_info_time = 0.0
print(f"[MusicMonitor] start api={api_url}", flush=True)
while True:
cfg = load_monitor_config()
if interval is not None:
cfg["interval_sec"] = interval
if cli_targets:
cfg["targets"] = cli_targets
if allow_all is not None:
cfg["allow_all"] = allow_all
sleep_sec = max(0.3, float(cfg.get("interval_sec", 1.0)))
holdover_sec = max(0, int(cfg.get("holdover_ms", 1500))) / 1000
keep_last = bool(cfg.get("keep_last_when_none", True))
try:
previous_cover_hash = ((state.get("current") or {}).get("cover_hash") or "")
info = await get_media_info(cfg, previous_cover_hash)
now = time.time()
resume_reason = await maybe_auto_resume(info, cfg) if info is not None else ""
if resume_reason:
info["playing"] = True
reason_text = "进度停滞" if resume_reason == "stalled" else "检测到暂停"
print(f"[MusicMonitor] auto resume playback ({reason_text})", flush=True)
if info is None:
monitor = state.setdefault("monitor", {})
monitor["online"] = False
monitor["updated_at"] = now
monitor["source"] = ""
monitor["platform"] = cfg.get("platform", "")
monitor["targets"] = cfg.get("targets", [])
monitor["allow_all"] = bool(cfg.get("allow_all", False))
within_holdover = bool(last_info_time) and now - last_info_time <= holdover_sec
if keep_last and within_holdover:
# SMTC 偶发丢一帧时短暂保留,避免页面闪烁。
pass
else:
# 超过保留时间后不能继续把旧歌曲伪装成“已暂停”。
# 仅清理播放器状态,保留点歌队列等独立数据。
empty = default_state()
state["playing"] = False
state["current"] = empty["current"]
if not keep_last:
state["playlist"] = []
else:
last_info_time = now
state["playing"] = info["playing"]
prev_cover = (state.get("current") or {}).get("cover", "")
prev_cover_hash = (state.get("current") or {}).get("cover_hash", "")
# 封面: 新值优先, 空值保留旧值(防闪)
new_cover = info.get("cover", "") or prev_cover
new_cover_hash = info.get("cover_hash", "") or prev_cover_hash
# 如果异步封面任务刚完成,本轮 info 可能尚未带 cover;这里再查一次缓存。
if not new_cover:
cached = _COVER_CACHE.get(cover_cache_key(info.get("title", ""), info.get("artist", "")))
if cached:
new_cover, new_cover_hash = cached
state["current"] = {
"title": info["title"],
"artist": info["artist"],
"cover": new_cover,
"cover_hash": new_cover_hash,
"duration": info["duration"],
"progress": info["progress"],
"source": info["source"],
}
state["monitor"] = {
"online": True,
"source": info["source"],
"platform": cfg.get("platform", ""),
"updated_at": info["updated_at"],
"targets": cfg.get("targets", []),
"allow_all": bool(cfg.get("allow_all", False)),
}
ok = await post_state(state, api_url)
if not ok:
await save_state(state)
cur = state.get("current") or {}
print(f"[MusicMonitor] {cur.get('title', '')} - {cur.get('artist', '')} "
f"({cur.get('progress', 0):.0f}s/{cur.get('duration', 0):.0f}s) "
f"playing={state.get('playing')} source={cur.get('source', '')}", flush=True)
except Exception as e:
print(f"[MusicMonitor] error: {e}", flush=True)
traceback.print_exc()
await asyncio.sleep(sleep_sec)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="BGI 直播间音乐状态同步")
parser.add_argument("--port", type=int, default=8086, help="本地 Web 服务端口号")
parser.add_argument("--interval", type=float, default=None, help="轮询间隔(秒),默认读取 config/config.json")
parser.add_argument("--allow-all", action="store_true", help="允许同步任意媒体会话(适合网页版播放器)")
parser.add_argument("--filter", default="", help="额外过滤关键字,匹配 source_app_display_name/aumid")
args = parser.parse_args()
targets = None
if args.filter:
cfg = load_monitor_config()
targets = list(cfg.get("targets", [])) + [args.filter]
asyncio.run(run_monitor(args.port, args.interval, targets, True if args.allow_all else None))
+320
View File
@@ -0,0 +1,320 @@
"""网易云音乐二维码登录会话。"""
from __future__ import annotations
import asyncio
import base64
import http.cookiejar
import http.cookies
import io
import json
import logging
import secrets
import string
import time
import urllib.parse
import urllib.request
from typing import Any, Callable
from Crypto.Cipher import AES
try:
from .netease_resolver import NeteaseResolver
except ImportError:
from netease_resolver import NeteaseResolver
_USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
_WEAPI_NONCE = b"0CoJUm6Qyw8W8jud"
_WEAPI_IV = b"0102030405060708"
_WEAPI_PUBLIC_EXPONENT = 0x10001
_WEAPI_MODULUS = int(
"00e0b509f6259df8642dbc35662901477df22677ec152b5ff68ace615bb7b725"
"152b3ab17a876aea8a5aa76d2e417629ec4ee341f56135fccf695280104e0312"
"ecbda92557c93870114af6c9d05c4f7f0c3685b7a46bee255932575cce10b424"
"d813cfe4875d3e82047b97ddef52741d546b8e289dc6935b3ece0462db0a22b8"
"e7",
16,
)
_SECRET_ALPHABET = string.ascii_letters + string.digits
def _aes_encrypt_base64(content: bytes, key: bytes) -> bytes:
padding = AES.block_size - (len(content) % AES.block_size)
padded = content + bytes([padding]) * padding
encrypted = AES.new(key, AES.MODE_CBC, _WEAPI_IV).encrypt(padded)
return base64.b64encode(encrypted)
def _weapi_form(data: dict[str, Any], *, secret_key: str | None = None) -> dict[str, str]:
secret = secret_key or "".join(secrets.choice(_SECRET_ALPHABET) for _ in range(16))
if len(secret.encode("ascii")) != 16:
raise ValueError("weapi secret key must be 16 ASCII bytes")
serialized = json.dumps(data, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
first_pass = _aes_encrypt_base64(serialized, _WEAPI_NONCE)
params = _aes_encrypt_base64(first_pass, secret.encode("ascii")).decode("ascii")
reversed_secret = int.from_bytes(secret[::-1].encode("ascii"), "big")
enc_sec_key = format(
pow(reversed_secret, _WEAPI_PUBLIC_EXPONENT, _WEAPI_MODULUS),
"x",
).zfill(256)
return {"params": params, "encSecKey": enc_sec_key}
def _render_qr_png(content: str) -> bytes:
try:
import qrcode
from qrcode.constants import ERROR_CORRECT_M
except ImportError as exc:
raise RuntimeError("缺少 qrcode 依赖,请重新安装 requirements.txt") from exc
qr = qrcode.QRCode(
version=None,
error_correction=ERROR_CORRECT_M,
box_size=8,
border=3,
)
qr.add_data(content)
qr.make(fit=True)
image = qr.make_image(fill_color="black", back_color="white")
output = io.BytesIO()
image.save(output, format="PNG")
return output.getvalue()
def _request_json(
url: str,
*,
opener: urllib.request.OpenerDirector | None = None,
data: dict[str, Any] | None = None,
) -> tuple[dict[str, Any], Any]:
encoded_data = None
if data is not None:
encoded_data = urllib.parse.urlencode(data).encode("utf-8")
request = urllib.request.Request(
url,
data=encoded_data,
headers={
"User-Agent": _USER_AGENT,
"Referer": "https://music.163.com/",
"Content-Type": "application/x-www-form-urlencoded",
},
)
response = (opener or urllib.request.build_opener()).open(request, timeout=15)
payload = json.loads(response.read().decode("utf-8", errors="replace"))
if not isinstance(payload, dict):
raise RuntimeError("网易云登录接口返回格式无效")
return payload, response
def _request_weapi_json(
url: str,
data: dict[str, Any],
*,
opener: urllib.request.OpenerDirector,
secret_key: str | None = None,
) -> tuple[dict[str, Any], Any]:
return _request_json(
url,
opener=opener,
data=_weapi_form(data, secret_key=secret_key),
)
def _cookie_values(jar: http.cookiejar.CookieJar, response: Any, payload: dict[str, Any]) -> dict[str, str]:
values = {cookie.name: cookie.value for cookie in jar}
headers = getattr(response, "headers", None)
raw_headers = headers.get_all("Set-Cookie", []) if headers and hasattr(headers, "get_all") else []
for raw_header in raw_headers:
parsed = http.cookies.SimpleCookie()
try:
parsed.load(raw_header)
except http.cookies.CookieError:
continue
values.update({name: morsel.value for name, morsel in parsed.items()})
for raw_cookie in (
payload.get("cookie"),
(payload.get("data") or {}).get("cookie") if isinstance(payload.get("data"), dict) else None,
):
if not raw_cookie:
continue
parsed = http.cookies.SimpleCookie()
try:
parsed.load(str(raw_cookie))
except http.cookies.CookieError:
continue
values.update({name: morsel.value for name, morsel in parsed.items()})
return values
class NeteaseQrLogin:
"""服务端保存二维码 key,登录成功后只把 MUSIC_U 交给保存回调。"""
_STATUS_MESSAGES = {
"idle": "尚未开始扫码登录",
"awaiting_scan": "请使用网易云音乐客户端扫码",
"awaiting_confirm": "已扫码,请在手机上确认登录",
"completed": "登录成功,MUSIC_U 已自动保存",
"expired": "二维码已过期,请重新生成",
"failed": "网易云扫码登录失败",
}
def __init__(
self,
*,
api_base: Callable[[], str],
get_saved_music_u: Callable[[], str],
save_music_u: Callable[[str, dict[str, Any]], Any],
logger: logging.Logger,
account_checker: Callable[[str, str], Any] | None = None,
ttl_seconds: int = 180,
) -> None:
self.api_base = api_base
self.get_saved_music_u = get_saved_music_u
self.save_music_u = save_music_u
self.logger = logger
self.account_checker = account_checker
self.ttl_seconds = max(60, int(ttl_seconds))
self._lock = asyncio.Lock()
self._session: dict[str, Any] | None = None
def _base_url(self) -> str:
return str(self.api_base() or "https://music.163.com").rstrip("/")
def snapshot(self) -> dict[str, Any]:
session = self._session or {}
state = str(session.get("state") or "idle")
expires_at = float(session.get("expires_at") or 0)
expires_in = max(0, int(expires_at - time.time())) if expires_at else 0
return {
"success": True,
"state": state,
"message": str(session.get("message") or self._STATUS_MESSAGES.get(state, "")),
"expires_in": expires_in,
"has_qr_image": state in {"awaiting_scan", "awaiting_confirm"} and expires_in > 0,
"credential_configured": bool(str(self.get_saved_music_u() or "").strip()),
"account": session.get("account"),
}
def _start_sync(self) -> dict[str, Any]:
jar = http.cookiejar.CookieJar()
opener = urllib.request.build_opener(urllib.request.HTTPCookieProcessor(jar))
payload, _ = _request_weapi_json(
f"{self._base_url()}/weapi/login/qrcode/unikey",
{"type": 1, "csrf_token": ""},
opener=opener,
)
key = str(payload.get("unikey") or (payload.get("data") or {}).get("unikey") or "").strip()
if int(payload.get("code") or 0) != 200 or not key:
raise RuntimeError(f"网易云二维码申请失败 code={payload.get('code')}")
now = time.time()
return {
"state": "awaiting_scan",
"message": self._STATUS_MESSAGES["awaiting_scan"],
"created_at": now,
"expires_at": now + self.ttl_seconds,
"key": key,
# The NetEase client currently accepts the exact HTTP URL emitted by
# the official web login page. Using HTTPS shows an unsupported-login warning.
"qr_url": f"http://music.163.com/login?codekey={key}",
"jar": jar,
"opener": opener,
"account": None,
}
async def start(self) -> dict[str, Any]:
async with self._lock:
try:
self._session = await asyncio.to_thread(self._start_sync)
except Exception as exc:
self._session = {
"state": "failed",
"message": str(exc) or type(exc).__name__,
"expires_at": 0,
}
self.logger.warning("[网易云扫码登录] 二维码申请失败: %s", type(exc).__name__)
return self.snapshot()
self.logger.info("[网易云扫码登录] 二维码已生成,等待扫码")
return self.snapshot()
def _poll_sync(self, session: dict[str, Any]) -> dict[str, Any]:
url = f"{self._base_url()}/weapi/login/qrcode/client/login"
payload, response = _request_weapi_json(
url,
{"key": session["key"], "type": 1, "csrf_token": ""},
opener=session["opener"],
)
code = int(payload.get("code") or 0)
if code == 801:
return {"state": "awaiting_scan", "message": self._STATUS_MESSAGES["awaiting_scan"]}
if code == 802:
return {"state": "awaiting_confirm", "message": self._STATUS_MESSAGES["awaiting_confirm"]}
if code == 800:
return {"state": "expired", "message": self._STATUS_MESSAGES["expired"]}
if code != 803:
return {
"state": "failed",
"message": str(payload.get("message") or f"网易云扫码状态异常 code={code}"),
}
music_u = str(_cookie_values(session["jar"], response, payload).get("MUSIC_U") or "").strip()
if not music_u:
return {"state": "failed", "message": "扫码成功但响应中缺少 MUSIC_U"}
return {"state": "authenticated", "music_u": music_u}
async def _check_account(self, music_u: str) -> dict[str, Any]:
if self.account_checker:
result = self.account_checker(self._base_url(), music_u)
return await result if asyncio.iscoroutine(result) else result
resolver = NeteaseResolver(self.logger, api_base=self._base_url(), music_u=music_u)
return await resolver.account_status()
async def poll(self) -> dict[str, Any]:
async with self._lock:
if not self._session:
return self.snapshot()
state = str(self._session.get("state") or "idle")
if state in {"completed", "expired", "failed"}:
return self.snapshot()
if time.time() >= float(self._session.get("expires_at") or 0):
self._session.update(state="expired", message=self._STATUS_MESSAGES["expired"])
return self.snapshot()
try:
result = await asyncio.to_thread(self._poll_sync, self._session)
if result.get("state") != "authenticated":
self._session.update(result)
return self.snapshot()
music_u = str(result.get("music_u") or "")
account = await self._check_account(music_u)
if not account.get("authenticated"):
self._session.update(state="failed", message="扫码 Cookie 登录验证失败")
return self.snapshot()
callback_result = self.save_music_u(music_u, account)
if asyncio.iscoroutine(callback_result):
await callback_result
self._session.update(
state="completed",
message=self._STATUS_MESSAGES["completed"],
account={
"user_id": str(account.get("user_id") or ""),
"nickname": str(account.get("nickname") or "网易云用户"),
"vip_type": int(account.get("vip_type") or 0),
},
)
self.logger.info("[网易云扫码登录] 登录成功,MUSIC_U 已自动保存")
return self.snapshot()
except Exception as exc:
self._session.update(
state="failed",
message=str(exc) or f"扫码状态查询异常: {type(exc).__name__}",
)
self.logger.warning("[网易云扫码登录] 状态查询或保存失败: %s", type(exc).__name__)
return self.snapshot()
async def qr_png(self) -> bytes:
async with self._lock:
if not self._session or self._session.get("state") not in {"awaiting_scan", "awaiting_confirm"}:
raise RuntimeError("当前没有可用的网易云登录二维码")
if time.time() >= float(self._session.get("expires_at") or 0):
self._session.update(state="expired", message=self._STATUS_MESSAGES["expired"])
raise RuntimeError("网易云登录二维码已过期")
content = str(self._session.get("qr_url") or "")
return await asyncio.to_thread(_render_qr_png, content)
+224
View File
@@ -0,0 +1,224 @@
"""Resolve NetEase song IDs to short-lived playable audio URLs."""
from __future__ import annotations
import logging
import re
import urllib.parse
from typing import Any
import aiohttp
class NeteaseResolver:
def __init__(
self,
logger: logging.Logger,
*,
api_base: str = "https://music.163.com",
music_u: str = "",
):
self.logger = logger
self.api_base = api_base.rstrip("/")
self.music_u = str(music_u or "").strip()
self.last_error_code = ""
def update_api_base(self, api_base: str):
self.api_base = str(api_base or "https://music.163.com").rstrip("/")
def update_auth(self, music_u: str):
self.music_u = str(music_u or "").strip()
def _headers(self) -> dict[str, str]:
cookie = "os=pc; appver=2.9.8;"
if self.music_u:
cookie += f" MUSIC_U={self.music_u};"
return {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36",
"Referer": "https://music.163.com/",
"Cookie": cookie,
"X-Real-IP": "218.75.111.114",
"X-Forwarded-For": "218.75.111.114",
}
async def _probe_url(self, url: str) -> str:
timeout = aiohttp.ClientTimeout(total=12)
async with aiohttp.ClientSession(headers=self._headers(), timeout=timeout) as session:
try:
async with session.get(url, allow_redirects=True, headers={"Range": "bytes=0-1"}) as resp:
content_type = str(resp.headers.get("Content-Type", "")).lower()
if resp.status in (200, 206) and ("audio" in content_type or "octet-stream" in content_type):
return str(resp.url)
except Exception:
return ""
return ""
async def _get_json(self, url: str) -> dict[str, Any] | None:
timeout = aiohttp.ClientTimeout(total=12)
try:
async with aiohttp.ClientSession(headers=self._headers(), timeout=timeout) as session:
async with session.get(url, allow_redirects=True) as resp:
if resp.status != 200:
return None
payload = await resp.json(content_type=None)
return payload if isinstance(payload, dict) else None
except Exception:
return None
async def account_status(self) -> dict[str, Any]:
if not self.music_u:
return {"authenticated": False}
payload = await self._get_json(f"{self.api_base}/api/nuser/account/get")
if not payload:
return {"authenticated": False}
account = payload.get("account")
profile = payload.get("profile")
if not isinstance(account, dict) or not account.get("id"):
return {"authenticated": False}
profile = profile if isinstance(profile, dict) else {}
return {
"authenticated": True,
"user_id": str(account.get("id") or ""),
"nickname": str(profile.get("nickname") or "网易云用户"),
"vip_type": int(account.get("vipType") or profile.get("vipType") or 0),
}
async def _fetch_player_entry(self, song_id: str) -> dict[str, Any] | None:
url = (
f"{self.api_base}/api/song/enhance/player/url"
f"?ids=%5B{urllib.parse.quote(song_id)}%5D&br=320000"
)
payload = await self._get_json(url)
entries = payload.get("data") if payload else None
if not isinstance(entries, list) or not entries or not isinstance(entries[0], dict):
return None
return entries[0]
@staticmethod
def _is_trial_entry(entry: dict[str, Any]) -> bool:
if entry.get("freeTrialInfo"):
return True
privilege = entry.get("freeTrialPrivilege")
if not isinstance(privilege, dict):
return False
return any(bool(privilege.get(key)) for key in ("resConsumable", "userConsumable", "listenType"))
@staticmethod
def playlist_id(value: str | int) -> str:
match = re.search(r"(?:playlist\?id=|\bid=)?(\d{5,})", str(value or ""))
return match.group(1) if match else ""
async def fetch_playlist(self, playlist: str | int) -> dict[str, Any] | None:
playlist_id = self.playlist_id(playlist)
if not playlist_id:
return None
timeout = aiohttp.ClientTimeout(total=15)
url = f"{self.api_base}/api/playlist/detail?id={urllib.parse.quote(playlist_id)}"
try:
async with aiohttp.ClientSession(headers=self._headers(), timeout=timeout) as session:
async with session.get(url, allow_redirects=True) as resp:
if resp.status != 200:
return None
payload = await resp.json(content_type=None)
except Exception as exc:
self.logger.warning(f"[mpv] 获取网易云歌单失败: {exc}")
return None
result = payload.get("result") if isinstance(payload, dict) else None
if not isinstance(result, dict):
return None
tracks = result.get("tracks")
if not isinstance(tracks, list):
return None
songs: list[dict[str, Any]] = []
for track in tracks:
if not isinstance(track, dict):
continue
song_id = self.playlist_id(track.get("id", ""))
if not song_id:
continue
artists_raw = track.get("artists") or track.get("ar") or []
artists = "/".join(
str(item.get("name") or "").strip()
for item in artists_raw
if isinstance(item, dict) and str(item.get("name") or "").strip()
)
album = track.get("album") or track.get("al") or {}
duration_ms = track.get("duration") or track.get("dt") or 0
try:
duration_sec = max(0, int(duration_ms) // 1000)
except (TypeError, ValueError):
duration_sec = 0
songs.append({
"id": song_id,
"name": str(track.get("name") or f"歌曲{song_id}"),
"artist": artists or "未知歌手",
"duration_sec": duration_sec,
"cover": str(album.get("picUrl") or "") if isinstance(album, dict) else "",
"source": "background_playlist",
"playlist_id": playlist_id,
})
if not songs:
return None
return {
"id": playlist_id,
"name": str(result.get("name") or f"歌单{playlist_id}"),
"songs": songs,
}
async def fetch_song_detail(self, song_id: str | int) -> dict[str, Any] | None:
clean_id = self.playlist_id(song_id)
if not clean_id:
return None
timeout = aiohttp.ClientTimeout(total=12)
url = f"{self.api_base}/api/song/detail/?id={clean_id}&ids=[{clean_id}]"
try:
async with aiohttp.ClientSession(headers=self._headers(), timeout=timeout) as session:
async with session.get(url, allow_redirects=True) as resp:
payload = await resp.json(content_type=None)
except Exception:
return None
songs = payload.get("songs") if isinstance(payload, dict) else None
if not isinstance(songs, list) or not songs or not isinstance(songs[0], dict):
return None
track = songs[0]
album = track.get("album") or track.get("al") or {}
return {
"cover": str(album.get("picUrl") or album.get("blurPicUrl") or "") if isinstance(album, dict) else "",
}
async def resolve(self, song: dict[str, Any]) -> dict[str, Any] | None:
self.last_error_code = ""
song_id = "".join(ch for ch in str(song.get("id", "")) if ch.isdigit())
if not song_id:
self.last_error_code = "invalid_song_id"
return None
entry = await self._fetch_player_entry(song_id)
if entry:
if self._is_trial_entry(entry):
reason = "登录已失效或账号没有完整播放权益" if self.music_u else "未配置网易云登录"
self.logger.warning(
f"[mpv] 拒绝播放试听片段: {song.get('name')} ({song_id}){reason}"
)
self.last_error_code = "preview_only"
return None
player_url = str(entry.get("url") or "").strip()
if player_url:
resolved = await self._probe_url(player_url)
if resolved:
return {
"url": resolved,
"source": "netease.player.auth" if self.music_u else "netease.player",
"song_id": song_id,
}
# NetEase's public outer URL provides a short-lived CDN redirect for songs available to the current region/account tier.
outer = f"{self.api_base}/song/media/outer/url?id={urllib.parse.quote(song_id)}.mp3"
resolved = await self._probe_url(outer)
if not resolved:
self.last_error_code = "unavailable"
self.logger.warning(f"[mpv] 无法获取可播放地址: {song.get('name')} ({song_id})")
return None
return {
"url": resolved,
"source": "netease.outer",
"song_id": song_id,
}
+419
View File
@@ -0,0 +1,419 @@
"""兑换码的 SQLite 存储与原子领取逻辑。"""
from __future__ import annotations
import asyncio
import sqlite3
import threading
from contextlib import contextmanager
from datetime import UTC, datetime, timedelta, timezone
from pathlib import Path
from typing import Any
BEIJING_TZ = timezone(timedelta(hours=8), name="Asia/Shanghai")
def normalize_code(value: str) -> str:
return str(value or "").strip().casefold()
def utc_now() -> str:
return datetime.now(UTC).isoformat(timespec="milliseconds").replace("+00:00", "Z")
def parse_beijing_datetime(value: str) -> str:
raw = str(value or "").strip()
if not raw:
raise ValueError("生效时间和失效时间不能为空")
try:
parsed = datetime.fromisoformat(raw.replace("Z", "+00:00"))
except ValueError as exc:
raise ValueError("时间格式无效") from exc
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=BEIJING_TZ)
return parsed.astimezone(UTC).isoformat(timespec="milliseconds").replace("+00:00", "Z")
def to_beijing_datetime(value: str | None) -> str:
if not value:
return ""
parsed = datetime.fromisoformat(str(value).replace("Z", "+00:00"))
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=UTC)
return parsed.astimezone(BEIJING_TZ).isoformat(timespec="minutes")
_SCHEMA = """
CREATE TABLE IF NOT EXISTS redemption_codes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code_normalized TEXT NOT NULL UNIQUE,
code_display TEXT NOT NULL,
points INTEGER NOT NULL CHECK(points > 0),
starts_at_utc TEXT NOT NULL,
ends_at_utc TEXT NOT NULL,
max_redemptions INTEGER CHECK(max_redemptions IS NULL OR max_redemptions > 0),
redeemed_count INTEGER NOT NULL DEFAULT 0 CHECK(redeemed_count >= 0),
enabled INTEGER NOT NULL DEFAULT 1 CHECK(enabled IN (0, 1)),
deleted_at_utc TEXT,
created_at_utc TEXT NOT NULL,
updated_at_utc TEXT NOT NULL,
CHECK(starts_at_utc < ends_at_utc)
);
CREATE TABLE IF NOT EXISTS redemption_records (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code_id INTEGER NOT NULL REFERENCES redemption_codes(id),
code_display TEXT NOT NULL,
platform TEXT NOT NULL,
platform_user_id TEXT NOT NULL,
display_name TEXT NOT NULL,
points INTEGER NOT NULL,
balance_before INTEGER NOT NULL,
balance_after INTEGER,
status TEXT NOT NULL DEFAULT 'pending',
redeemed_at_utc TEXT NOT NULL,
completed_at_utc TEXT,
UNIQUE(code_id, platform, platform_user_id)
);
CREATE INDEX IF NOT EXISTS idx_redemption_codes_active
ON redemption_codes(deleted_at_utc, enabled, starts_at_utc, ends_at_utc);
CREATE INDEX IF NOT EXISTS idx_redemption_records_code_time
ON redemption_records(code_id, redeemed_at_utc DESC);
CREATE INDEX IF NOT EXISTS idx_redemption_records_user_time
ON redemption_records(platform, platform_user_id, redeemed_at_utc DESC);
"""
class RedemptionCodeStore:
"""用短事务处理兑换码,避免并发超领和重复领取。
当提供 stats_store 时,所有读写都复用 stats_store 的唯一连接与串行写队列,
消除多连接写同一文件导致的 "database is locked" 争用。
"""
def __init__(self, database_path: str | Path, stats_store: Any | None = None):
self.database_path = Path(database_path)
self._store = stats_store
self._lock = threading.RLock()
self.database_path.parent.mkdir(parents=True, exist_ok=True)
with self._connect() as connection:
connection.executescript(_SCHEMA)
connection.commit()
@contextmanager
def _connect(self):
connection = sqlite3.connect(self.database_path, timeout=5.0)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA journal_mode=WAL")
connection.execute("PRAGMA synchronous=NORMAL")
connection.execute("PRAGMA foreign_keys=ON")
connection.execute("PRAGMA busy_timeout=5000")
try:
yield connection
finally:
connection.close()
def _run(self, func, *args):
"""在 stats_store 单连接(或自有连接)上执行 func(connection, *args)。"""
if self._store is not None:
return self._store.execute(func, *args)
return asyncio.to_thread(self._run_direct, func, args)
def _run_direct(self, func, args):
with self._lock, self._connect() as connection:
return func(connection, *args)
@staticmethod
def _serialize_code(row: sqlite3.Row) -> dict[str, Any]:
max_redemptions = row["max_redemptions"]
redeemed_count = int(row["redeemed_count"] or 0)
return {
"id": int(row["id"]),
"code": row["code_display"],
"points": int(row["points"]),
"starts_at": to_beijing_datetime(row["starts_at_utc"]),
"ends_at": to_beijing_datetime(row["ends_at_utc"]),
"max_redemptions": int(max_redemptions) if max_redemptions is not None else None,
"redeemed_count": redeemed_count,
"remaining_count": max(0, int(max_redemptions) - redeemed_count) if max_redemptions is not None else None,
"enabled": bool(row["enabled"]),
"created_at": to_beijing_datetime(row["created_at_utc"]),
"updated_at": to_beijing_datetime(row["updated_at_utc"]),
}
@staticmethod
def _serialize_record(row: sqlite3.Row) -> dict[str, Any]:
return {
"id": int(row["id"]),
"code_id": int(row["code_id"]),
"code": row["code_display"],
"platform": row["platform"],
"uid": row["platform_user_id"],
"uname": row["display_name"],
"points": int(row["points"]),
"balance_before": int(row["balance_before"]),
"balance_after": int(row["balance_after"]) if row["balance_after"] is not None else None,
"status": row["status"],
"redeemed_at": to_beijing_datetime(row["redeemed_at_utc"]),
"completed_at": to_beijing_datetime(row["completed_at_utc"]),
}
async def create_code(
self,
*,
code: str,
points: int,
starts_at: str,
ends_at: str,
max_redemptions: int | None,
enabled: bool = True,
) -> dict[str, Any]:
return await self._run(
self._create_code,
code,
points,
starts_at,
ends_at,
max_redemptions,
enabled,
)
def _create_code(
self,
connection,
code: str,
points: int,
starts_at: str,
ends_at: str,
max_redemptions: int | None,
enabled: bool,
) -> dict[str, Any]:
display = str(code or "").strip()
normalized = normalize_code(display)
if not normalized:
raise ValueError("兑换码不能为空")
if len(display) > 100:
raise ValueError("兑换码不能超过100个字符")
points = int(points)
if points <= 0:
raise ValueError("兑换积分必须是正整数")
if max_redemptions in ("", None):
max_value = None
else:
max_value = int(max_redemptions)
if max_value <= 0:
raise ValueError("总兑换次数必须是正整数,或留空表示不限")
starts_utc = parse_beijing_datetime(starts_at)
ends_utc = parse_beijing_datetime(ends_at)
if starts_utc >= ends_utc:
raise ValueError("失效时间必须晚于生效时间")
now = utc_now()
try:
cursor = connection.execute(
"""INSERT INTO redemption_codes (
code_normalized, code_display, points, starts_at_utc, ends_at_utc,
max_redemptions, enabled, created_at_utc, updated_at_utc
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(normalized, display, points, starts_utc, ends_utc, max_value, int(bool(enabled)), now, now),
)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
if "code_normalized" in str(exc) or "UNIQUE constraint" in str(exc):
raise ValueError("兑换码已存在") from exc
raise
row = connection.execute("SELECT * FROM redemption_codes WHERE id=?", (cursor.lastrowid,)).fetchone()
assert row is not None
return self._serialize_code(row)
async def list_codes(self) -> list[dict[str, Any]]:
return await self._run(self._list_codes)
def _list_codes(self, connection) -> list[dict[str, Any]]:
rows = connection.execute(
"SELECT * FROM redemption_codes WHERE deleted_at_utc IS NULL ORDER BY id DESC"
).fetchall()
return [self._serialize_code(row) for row in rows]
async def list_records(self, *, code_id: int | None = None, limit: int = 500) -> list[dict[str, Any]]:
return await self._run(self._list_records, code_id, limit)
def _list_records(self, connection, code_id: int | None, limit: int) -> list[dict[str, Any]]:
limit = max(1, min(2000, int(limit)))
if code_id is None:
rows = connection.execute(
"SELECT * FROM redemption_records ORDER BY id DESC LIMIT ?", (limit,)
).fetchall()
else:
rows = connection.execute(
"SELECT * FROM redemption_records WHERE code_id=? ORDER BY id DESC LIMIT ?",
(int(code_id), limit),
).fetchall()
return [self._serialize_record(row) for row in rows]
async def set_enabled(self, code_id: int, enabled: bool) -> dict[str, Any]:
return await self._run(self._set_enabled, code_id, enabled)
def _set_enabled(self, connection, code_id: int, enabled: bool) -> dict[str, Any]:
cursor = connection.execute(
"UPDATE redemption_codes SET enabled=?, updated_at_utc=? WHERE id=? AND deleted_at_utc IS NULL",
(int(bool(enabled)), utc_now(), int(code_id)),
)
if cursor.rowcount != 1:
connection.rollback()
raise ValueError("兑换码不存在")
connection.commit()
row = connection.execute("SELECT * FROM redemption_codes WHERE id=?", (int(code_id),)).fetchone()
assert row is not None
return self._serialize_code(row)
async def delete_code(self, code_id: int) -> None:
await self._run(self._delete_code, code_id)
def _delete_code(self, connection, code_id: int) -> None:
now = utc_now()
cursor = connection.execute(
"""UPDATE redemption_codes
SET enabled=0, deleted_at_utc=?, updated_at_utc=?
WHERE id=? AND deleted_at_utc IS NULL""",
(now, now, int(code_id)),
)
if cursor.rowcount != 1:
connection.rollback()
raise ValueError("兑换码不存在")
connection.commit()
async def reserve(
self,
message: str,
*,
platform: str,
platform_user_id: str,
display_name: str,
balance_before: int,
) -> dict[str, Any]:
return await self._run(
self._reserve,
message,
platform,
platform_user_id,
display_name,
balance_before,
)
def _reserve(
self,
connection,
message: str,
platform: str,
platform_user_id: str,
display_name: str,
balance_before: int,
) -> dict[str, Any]:
normalized = normalize_code(message)
if not normalized:
return {"status": "unknown"}
now = utc_now()
connection.execute("BEGIN IMMEDIATE")
try:
code = connection.execute(
"SELECT * FROM redemption_codes WHERE code_normalized=? AND deleted_at_utc IS NULL",
(normalized,),
).fetchone()
if code is None:
connection.rollback()
return {"status": "unknown"}
if not bool(code["enabled"]):
connection.rollback()
return {"status": "disabled"}
if now < code["starts_at_utc"]:
connection.rollback()
return {"status": "not_started"}
if now >= code["ends_at_utc"]:
connection.rollback()
return {"status": "expired"}
existing = connection.execute(
"""SELECT status FROM redemption_records
WHERE code_id=? AND platform=? AND platform_user_id=?""",
(code["id"], platform, platform_user_id),
).fetchone()
if existing is not None:
connection.rollback()
return {"status": "already_redeemed"}
max_redemptions = code["max_redemptions"]
if max_redemptions is not None and int(code["redeemed_count"]) >= int(max_redemptions):
connection.rollback()
return {"status": "exhausted"}
cursor = connection.execute(
"""INSERT INTO redemption_records (
code_id, code_display, platform, platform_user_id, display_name,
points, balance_before, status, redeemed_at_utc
) VALUES (?, ?, ?, ?, ?, ?, ?, 'pending', ?)""",
(
code["id"], code["code_display"], platform, platform_user_id,
display_name, code["points"], int(balance_before), now,
),
)
updated = connection.execute(
"""UPDATE redemption_codes
SET redeemed_count=redeemed_count+1, updated_at_utc=?
WHERE id=? AND (max_redemptions IS NULL OR redeemed_count < max_redemptions)""",
(now, code["id"]),
)
if updated.rowcount != 1:
connection.rollback()
return {"status": "exhausted"}
connection.commit()
return {
"status": "reserved",
"record_id": int(cursor.lastrowid),
"code_id": int(code["id"]),
"code": code["code_display"],
"points": int(code["points"]),
}
except sqlite3.IntegrityError:
connection.rollback()
return {"status": "already_redeemed"}
except Exception:
connection.rollback()
raise
async def finalize(self, record_id: int, balance_after: int) -> None:
await self._run(self._finalize, record_id, balance_after)
def _finalize(self, connection, record_id: int, balance_after: int) -> None:
cursor = connection.execute(
"""UPDATE redemption_records
SET status='completed', balance_after=?, completed_at_utc=?
WHERE id=? AND status='pending'""",
(int(balance_after), utc_now(), int(record_id)),
)
if cursor.rowcount != 1:
connection.rollback()
raise ValueError("兑换记录不存在或已经完成")
connection.commit()
async def cancel(self, record_id: int) -> None:
await self._run(self._cancel, record_id)
def _cancel(self, connection, record_id: int) -> None:
connection.execute("BEGIN IMMEDIATE")
try:
row = connection.execute(
"SELECT code_id FROM redemption_records WHERE id=? AND status='pending'",
(int(record_id),),
).fetchone()
if row is None:
connection.rollback()
return
connection.execute("DELETE FROM redemption_records WHERE id=?", (int(record_id),))
connection.execute(
"""UPDATE redemption_codes
SET redeemed_count=MAX(0, redeemed_count-1), updated_at_utc=?
WHERE id=?""",
(utc_now(), int(row["code_id"])),
)
connection.commit()
except Exception:
connection.rollback()
raise
+1733
View File
File diff suppressed because it is too large Load Diff
+127
View File
@@ -0,0 +1,127 @@
"""TTS 状态监控窗口
独立小黑窗显示 TTS 引擎状态、最近合成文本和耗时。
启动参数:
tts_monitor.py --state-file data/tts_state.json
"""
import argparse
import json
import os
import sys
import time
from pathlib import Path
APP_DIR = Path(__file__).resolve().parent
if str(APP_DIR) not in sys.path:
sys.path.insert(0, str(APP_DIR))
from core.runtime_paths import DATA_DIR, ensure_runtime_dirs
def clear():
os.system("cls" if os.name == "nt" else "clear")
def default_state() -> dict:
return {
"enabled": False,
"provider": "none",
"model_loaded": False,
"model_name": "",
"last_text": "",
"last_duration_ms": 0,
"last_error": "",
"total_synthesized": 0,
"total_errors": 0,
"recent_events": [],
"updated_at": 0,
}
def load_state(state_file: Path) -> dict:
if not state_file.exists():
return default_state()
try:
data = json.loads(state_file.read_text(encoding="utf-8"))
if isinstance(data, dict):
state = default_state()
state.update(data)
return state
except Exception:
pass
return default_state()
def fmt_duration(ms: int) -> str:
if ms <= 0:
return "-"
if ms < 1000:
return f"{ms}ms"
return f"{ms/1000:.2f}s"
def render(state: dict):
clear()
print("=" * 60)
print(" TTS 状态监控")
print("=" * 60)
print()
enabled = state.get("enabled", False)
provider = state.get("provider", "none")
model_loaded = state.get("model_loaded", False)
print(f" TTS 启用: {'' if enabled else ''}")
print(f" 引擎: {provider}")
print(f" 模型加载: {'完成' if model_loaded else '未加载/加载中'}")
print(f" 模型名: {state.get('model_name', '') or '-'}")
print()
print(f" 累计合成: {state.get('total_synthesized', 0)}")
print(f" 累计失败: {state.get('total_errors', 0)}")
print(f" 上次合成: {state.get('last_text', '') or '-'}")
print(f" 合成耗时: {fmt_duration(state.get('last_duration_ms', 0))}")
print()
last_error = state.get("last_error", "")
if last_error:
print(f" [错误] {last_error}")
print()
print("-" * 60)
print(" 最近事件")
print("-" * 60)
events = state.get("recent_events", [])
if not events:
print(" (无)")
else:
for ev in events[-8:]:
ts = ev.get("time", "")
msg = ev.get("msg", "")
print(f" {ts} {msg}")
print()
print(" 按 Ctrl+C 关闭本窗口")
def main():
ensure_runtime_dirs()
parser = argparse.ArgumentParser(description="TTS 状态监控")
parser.add_argument("--state-file", default=str(DATA_DIR / "tts_state.json"), help="TTS 状态文件路径")
args = parser.parse_args()
state_file = Path(args.state_file)
print("等待 TTS 状态更新...")
time.sleep(0.5)
try:
while True:
state = load_state(state_file)
render(state)
time.sleep(0.5)
except KeyboardInterrupt:
print("\nTTS 监控已关闭")
sys.exit(0)
if __name__ == "__main__":
main()