Files
creator-hub/browser_gateway/agent.py
T

376 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Gateway outbound agent: dial the platform, dispatch tasks, push events.
方向反转后网关只有一条自己发起的 WS 连接(/v1/agent,连接层 Bearer key 认证):
- 心跳线程周期上报 {type: heartbeat, version},平台据此维护 online/last_seen_at;
- 读循环接收 task(线程池执行,复用 HTTP 路由语义)与 event_ack/event_nack;
- result 帧回传路由结果;重复 task id 幂等拒绝,不重复执行有副作用的操作;
- 事件桥把 SubscriptionManager 的 pending deliveries 推为 event 帧,
平台事务保存成功后回 event_ack;未确认批次在 nack 或断线重连后重推。
协议与平台侧 internal/controlplane/api/gateway_channel.go 一一对应。
"""
from __future__ import annotations
import base64
import json
import logging
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from typing import Any, Callable
import websocket
LOG = logging.getLogger("creatorhub.gateway")
GW_VERSION = "1.0.0-outbound"
HEARTBEAT_INTERVAL = 20.0
RECONNECT_BASE = 1.0
RECONNECT_MAX = 30.0
WORKER_COUNT = 4
IDEMPOTENT_STATUS = 409
@dataclass
class AgentConfig:
platform_url: str
access_key: str
version: str = GW_VERSION
heartbeat_interval: float = HEARTBEAT_INTERVAL
reconnect_base: float = RECONNECT_BASE
reconnect_max: float = RECONNECT_MAX
worker_count: int = WORKER_COUNT
class OutboundAgent:
"""单网关出站连接的完整生命周期;run() 阻塞运行,stop() 置位退出。"""
def __init__(self, gateway: Any, config: AgentConfig) -> None:
self.gateway = gateway
self.config = config
self.connected = threading.Event()
self.last_error = ""
self.reconnect_attempts = 0
self._stop_requested = threading.Event()
self._write_lock = threading.Lock()
self._conn: websocket.WebSocket | None = None
self._sent_events: dict[str, dict[str, float]] = {}
self._sent_lock = threading.Lock()
self._seen_tasks: set[int] = set()
self._seen_tasks_lock = threading.Lock()
self._push_wakeup = threading.Event()
# ---- 生命周期 -------------------------------------------------------
def run(self) -> None:
backoff = self.config.reconnect_base
while not self._stop_requested.is_set():
try:
self.reconnect_attempts += 1
conn = self._connect()
except Exception as exc:
self.last_error = str(exc) or exc.__class__.__name__
LOG.warning("platform connection failed: %s", self.last_error)
if self._wait(backoff):
break
backoff = min(backoff * 2, self.config.reconnect_max)
continue
self.reconnect_attempts = 0
self.last_error = ""
self._serve(conn)
backoff = self.config.reconnect_base
if not self._stop_requested.is_set():
LOG.info("platform connection lost; reconnecting")
def stop(self) -> None:
self._stop_requested.set()
with self._write_lock:
conn = self._conn
if conn is not None:
try:
conn.close()
except Exception:
pass
def _wait(self, seconds: float) -> bool:
return self._stop_requested.wait(seconds)
def _connect(self) -> websocket.WebSocket:
url = self.config.platform_url.rstrip("/") + "/v1/agent"
if url.startswith("https://"):
url = "wss://" + url[len("https://"):]
elif url.startswith("http://"):
url = "ws://" + url[len("http://"):]
conn = websocket.create_connection(
url,
header=[f"Authorization: Bearer {self.config.access_key}"],
timeout=10,
)
# timeout 仅用于握手;读循环必须阻塞,否则空闲连接会在 recv 超时后假性断线。
conn.settimeout(None)
with self._write_lock:
self._conn = conn
return conn
# ---- 会话线程 -------------------------------------------------------
def _serve(self, conn: websocket.WebSocket) -> None:
self.connected.set()
with self._sent_lock:
self._sent_events.clear()
executor = ThreadPoolExecutor(max_workers=self.config.worker_count, thread_name_prefix="gw-task")
heartbeat = threading.Thread(target=self._heartbeat_loop, args=(conn,), daemon=True)
sender = threading.Thread(target=self._event_sender_loop, args=(conn,), daemon=True)
heartbeat.start()
sender.start()
failure = ""
try:
while not self._stop_requested.is_set():
raw = conn.recv()
if not isinstance(raw, str):
continue
try:
frame = json.loads(raw)
except ValueError:
failure = "gateway channel message is not JSON"
break
if not isinstance(frame, dict):
failure = "gateway channel message invalid"
break
kind = frame.get("type")
if kind == "task":
future = executor.submit(self._handle_task, conn, frame)
future.add_done_callback(self._log_task_failure)
elif kind == "event_ack":
self._handle_event_ack(frame)
elif kind == "event_nack":
self._handle_event_nack(frame)
else:
failure = f"gateway channel message type invalid: {kind!r}"
break
except Exception as exc:
if not self._stop_requested.is_set():
failure = str(exc) or exc.__class__.__name__
LOG.exception("gateway channel receive loop failed: %s", failure)
finally:
self.connected.clear()
executor.shutdown(wait=False, cancel_futures=True)
self._stop_threads((heartbeat, sender))
with self._write_lock:
if self._conn is conn:
self._conn = None
try:
conn.close()
except Exception:
pass
if failure:
self.last_error = failure
@staticmethod
def _stop_threads(threads: tuple[threading.Thread, ...]) -> None:
for thread in threads:
thread.join(timeout=2)
def _heartbeat_loop(self, conn: websocket.WebSocket) -> None:
while not self._stop_requested.wait(self.config.heartbeat_interval):
try:
self._send(conn, {"type": "heartbeat", "version": self.config.version})
except Exception:
try:
conn.close()
except Exception:
pass
return
# ---- 帧读写 ---------------------------------------------------------
def _send(self, conn: websocket.WebSocket, value: dict) -> None:
payload = json.dumps(value, ensure_ascii=False, separators=(",", ":"))
with self._write_lock:
conn.send(payload)
# ---- 任务分发 -------------------------------------------------------
def _handle_task(self, conn: websocket.WebSocket, frame: dict) -> None:
task_id = frame.get("id")
method = frame.get("method")
path = frame.get("path")
if not isinstance(task_id, int) or not isinstance(method, str) or not isinstance(path, str):
self.last_error = "gateway task frame invalid"
try:
conn.close()
except Exception:
pass
return
with self._seen_tasks_lock:
duplicate = task_id in self._seen_tasks
self._seen_tasks.add(task_id)
if duplicate:
self._send(conn, {
"type": "result",
"id": task_id,
"status": IDEMPOTENT_STATUS,
"body": "",
"error": "duplicate task id; previous execution retained",
})
return
status, body, error = self._execute_task(method, path, frame.get("payload"))
self._send(conn, {
"type": "result",
"id": task_id,
"status": status,
"body": base64.b64encode(body).decode("ascii"),
"error": error,
})
def _execute_task(self, method: str, path: str, payload: Any) -> tuple[int, bytes, str]:
from .server.http import RequestError, route_gateway_request
if payload is None:
body: dict = {}
elif isinstance(payload, dict):
body = payload
else:
return 400, b"", "task payload must be a JSON object"
try:
result = route_gateway_request(self.gateway, method, path, {}, body)
except RequestError as exc:
return exc.status, b"", str(exc)
except FileNotFoundError as exc:
return 404, b"", str(exc)
except ValueError as exc:
return 400, b"", str(exc)
except Exception as exc: # 路由内部故障必须回传给平台,不能吞掉
LOG.exception("gateway task %s %s failed", method, path)
return 500, b"", "gateway operation failed"
if result is None:
return 204, b"", ""
if isinstance(result, tuple):
status, value = result
return int(status), json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode("utf-8"), ""
return 200, json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode("utf-8"), ""
@staticmethod
def _log_task_failure(future: Any) -> None:
exc = future.exception()
if exc is not None:
LOG.exception("gateway task dispatch failed", exc_info=exc)
# ---- 事件桥 ---------------------------------------------------------
def _event_sender_loop(self, conn: websocket.WebSocket) -> None:
manager = self.gateway.subscriptions
while not self._stop_requested.is_set():
try:
with manager.changed:
version = manager.version
self._push_wakeup.clear()
self._push_pending(conn, version)
with manager.changed:
manager.changed.wait_for(
lambda: self._stop_requested.is_set()
or manager.version != version
or self._push_wakeup.is_set(),
timeout=5,
)
except Exception:
if self._stop_requested.is_set():
return
LOG.exception("event sender loop failed; closing connection")
try:
conn.close()
except Exception:
pass
return
@staticmethod
def _subscription_items(manager: Any) -> list[tuple[str, Any]]:
# 从注册表快照枚举活动订阅;_items 缺失(测试替身)时视为空。
items = getattr(manager, "_items", {})
if not isinstance(items, dict):
return []
lock = getattr(manager, "_lock", None)
if not isinstance(lock, type(threading.RLock())):
return list(items.items())
with lock:
return list(items.items())
def _push_pending(self, conn: websocket.WebSocket, version: int) -> None:
# 每轮全量对账 pending 与已推记录的差集:新事件推送一次;nack/重连清空
# 已推记录后,同一批次会自然重推;无新差集时为空操作。
from .platform.douyin import DouyinError
manager = self.gateway.subscriptions
with self._sent_lock:
sent = self._sent_events
for alias, item in self._subscription_items(manager):
try:
current = manager._get(alias)
error = item.error or (
"event subscription replaced or stopped"
if current is not item or item.stopped
else ""
)
except (DouyinError, Exception) as exc:
error = str(exc) or exc.__class__.__name__
if error:
self._send(conn, {"type": "error", "subscription": "", "alias": alias, "error": error})
sent.pop(alias, None)
continue
try:
already = sent.setdefault(alias, {})
pending = [d for d in item.pending() if d["delivery_id"] not in already]
except Exception as exc:
self._send(conn, {"type": "error", "subscription": "", "alias": alias, "error": str(exc)})
sent.pop(alias, None)
continue
if not pending:
continue
subscription = f"{alias}:{item.session_id}"
self._send(conn, {
"type": "event",
"event": {"subscription": subscription, "alias": alias, "deliveries": pending},
})
stamp = time.monotonic()
already.update({d["delivery_id"]: stamp for d in pending})
LOG.info("event pushed alias=%s subscription=%s count=%s", alias, subscription, len(pending))
def _handle_event_ack(self, frame: dict) -> None:
alias = frame.get("alias")
ids = frame.get("delivery_ids")
if not isinstance(alias, str) or not isinstance(ids, list):
self.last_error = "event acknowledgement invalid"
return
manager = self.gateway.subscriptions
with self._sent_lock:
already = self._sent_events.get(alias, {})
unknown = [i for i in ids if i not in already]
if unknown:
# 平台确认了网关没有记录的投递(重连后旧 ack 等),仅清理已知项。
LOG.warning("event ack references unknown deliveries alias=%s count=%s", alias, len(unknown))
known = [i for i in ids if i in already]
if not known:
return
manager.ack(alias, known)
with self._sent_lock:
already = self._sent_events.get(alias, {})
for delivery_id in known:
already.pop(delivery_id, None)
def _handle_event_nack(self, frame: dict) -> None:
alias = frame.get("alias")
if not isinstance(alias, str):
return
with self._sent_lock:
self._sent_events.pop(alias, None)
self._wake_sender()
def _wake_sender(self) -> None:
# Condition.wait_for 只在 notify/超时时求值谓词;清空已推记录后必须唤醒。
self._push_wakeup.set()
manager = self.gateway.subscriptions
changed = getattr(manager, "changed", None)
if isinstance(changed, threading.Condition):
with changed:
changed.notify_all()