"""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()