diff --git a/proxy/bridge.py b/proxy/bridge.py index a8f004a..2161acf 100644 --- a/proxy/bridge.py +++ b/proxy/bridge.py @@ -1,7 +1,6 @@ import asyncio import logging import struct -import random from typing import List, Optional from urllib.parse import urlencode @@ -181,21 +180,22 @@ async def _cfproxy_worker_fallback(reader, writer, relay_init, label, worker_domains = proxy_config.cfproxy_worker_domains if not worker_domains: return False - - random.shuffle(worker_domains) - for worker_domain in worker_domains: - ws = None if is_test_dc else await cf_worker_pool.get(dc, worker_domain, fallback_dst) - if ws: - log.info("[%s] DC%d%s -> CF worker pool hit for %s", - label, dc, media_tag, fallback_dst) - else: - query = urlencode({ - 'dst': fallback_dst, - 'dc': str(dc), - }) - path = f'/apiws?{query}' + pooled = None if is_test_dc else await cf_worker_pool.get( + dc, fallback_dst, worker_domains) + if pooled: + ws, worker_domain = pooled + log.info("[%s] DC%d%s -> CF worker pool hit via %s for %s", + label, dc, media_tag, worker_domain, fallback_dst) + else: + query = urlencode({ + 'dst': fallback_dst, + 'dc': str(dc), + }) + path = f'/apiws?{query}' + ws = None + for worker_domain in cf_worker_pool.available_domains(worker_domains): log.info("[%s] DC%d%s -> trying CF worker %s for %s", label, dc, media_tag, worker_domain, fallback_dst) @@ -203,17 +203,20 @@ async def _cfproxy_worker_fallback(reader, writer, relay_init, label, ws = await RawWebSocket.connect(worker_domain, worker_domain, timeout=10.0, path=path) except Exception as exc: + cf_worker_pool.report_failure(worker_domain, exc) log.warning("[%s] DC%d%s CF worker %s failed: %s", label, dc, media_tag, worker_domain, repr(exc)) continue - stats.connections_cfproxy += 1 - await ws.send(relay_init) - await bridge_ws_reencrypt(reader, writer, ws, label, ctx, - dc=dc, is_media=is_media, - splitter=None) - return True - return False + if ws is None: + return False + + stats.connections_cfproxy += 1 + await ws.send(relay_init) + await bridge_ws_reencrypt(reader, writer, ws, label, ctx, + dc=dc, is_media=is_media, + splitter=None) + return True async def _cfproxy_fallback(reader, writer, relay_init, label, diff --git a/proxy/pool.py b/proxy/pool.py index ac532fd..ee95850 100644 --- a/proxy/pool.py +++ b/proxy/pool.py @@ -1,5 +1,6 @@ import asyncio import logging +import random import time from collections import deque @@ -221,73 +222,105 @@ class _CfWorkerPool: PER_DC_LIMIT = 1 def __init__(self): - self._idle: Dict[Tuple[int, str], deque] = {} - self._refilling: Set[Tuple[int, str]] = set() + self._idle: Dict[int, deque] = {} + self._refilling: Set[int] = set() + self._exhausted_until: Dict[str, float] = {} - async def get(self, dc: int, worker_domain: str, fallback_dst: str) -> Optional[RawWebSocket]: + async def get(self, dc: int, fallback_dst: str, + worker_domains: List[str] + ) -> Optional[Tuple[RawWebSocket, str]]: now = time.monotonic() - key = (dc, worker_domain) - bucket = self._idle.get(key) + bucket = self._idle.get(dc) if bucket is None: bucket = deque() - self._idle[key] = bucket + self._idle[dc] = bucket while bucket: - ws, created = bucket.popleft() + ws, created, worker_domain = bucket.popleft() age = now - created if (age > self.WS_POOL_MAX_AGE or ws._closed or ws.writer.transport.is_closing()): asyncio.create_task(self._quiet_close(ws)) continue stats.cf_pool_hits += 1 - log.debug("CF worker pool hit DC%d (age=%.1fs, left=%d)", - dc, age, len(bucket)) - self._schedule_refill(key, fallback_dst) - return ws + log.debug( + "CF worker pool hit DC%d via %s (age=%.1fs, left=%d)", + dc, worker_domain, age, len(bucket)) + self._schedule_refill(dc, fallback_dst, worker_domains) + return ws, worker_domain stats.cf_pool_misses += 1 - self._schedule_refill(key, fallback_dst) return None - def _schedule_refill(self, key, fallback_dst): - if key in self._refilling: + def _schedule_refill(self, dc, fallback_dst, worker_domains): + if dc in self._refilling: return - self._refilling.add(key) - asyncio.create_task(self._refill(key, fallback_dst)) + self._refilling.add(dc) + asyncio.create_task(self._refill( + dc, fallback_dst, list(worker_domains))) - async def _refill(self, key, fallback_dst): - dc, worker_domain = key + async def _refill(self, dc, fallback_dst, worker_domains): try: - bucket = self._idle.setdefault(key, deque()) - needed = min(proxy_config.pool_size - len(bucket), self.PER_DC_LIMIT) + bucket = self._idle.setdefault(dc, deque()) + target_size = min(proxy_config.pool_size, self.PER_DC_LIMIT) + needed = target_size - len(bucket) if needed <= 0: return - tasks = [asyncio.create_task( - self._connect_one(worker_domain, fallback_dst, dc)) - for _ in range(needed)] - for t in tasks: - try: - ws = await t - if ws: - bucket.append((ws, time.monotonic())) - except Exception: - pass + + for _ in range(needed): + connected = await self._connect_one( + worker_domains, fallback_dst, dc) + if connected is None: + break + ws, worker_domain = connected + bucket.append((ws, time.monotonic(), worker_domain)) log.debug("CF worker pool refilled DC%d: %d ready", dc, len(bucket)) finally: - self._refilling.discard(key) + self._refilling.discard(dc) - async def _connect_one(self, worker_domain, fallback_dst, dc) -> Optional[RawWebSocket]: + async def _connect_one(self, worker_domains, fallback_dst, dc): query = urlencode({ 'dst': fallback_dst, 'dc': str(dc), }) path = f'/apiws?{query}' - try: - return await RawWebSocket.connect( - worker_domain, worker_domain, timeout=8, path=path) - except Exception: - return None + for worker_domain in self.available_domains(worker_domains): + try: + ws = await RawWebSocket.connect( + worker_domain, worker_domain, timeout=8, path=path) + return ws, worker_domain + except Exception as exc: + self.report_failure(worker_domain, exc) + return None + + def available_domains(self, worker_domains: List[str]) -> List[str]: + now = time.time() + domains = list() + for domain in worker_domains: + if domain in domains: + continue + exhausted_until = self._exhausted_until.get(domain, 0) + if exhausted_until > now: + continue + if exhausted_until: + self._exhausted_until.pop(domain, None) + domains.append(domain) + random.shuffle(domains) + return domains + + def report_failure(self, worker_domain: str, exc: Exception) -> None: + return # TODO: check status code after daily limit reached + if not isinstance(exc, WsHandshakeError) or exc.status_code != 429: + return + + now = time.time() + if self._exhausted_until.get(worker_domain, 0) > now: + return + exhausted_until = now + (86400 - (now % 86400)) + self._exhausted_until[worker_domain] = exhausted_until + log.warning( + "CF worker %s reached its request limit, disabled for %d seconds", worker_domain, int(exhausted_until - now)) async def _quiet_close(self, ws): try: @@ -304,15 +337,16 @@ class _CfWorkerPool: if not cf_fallbacks or not proxy_config.cfproxy_worker_domains: return - for worker_domain in proxy_config.cfproxy_worker_domains: - for dc, fallback_dst in cf_fallbacks.items(): - self._schedule_refill((dc, worker_domain), fallback_dst) + worker_domains = list(proxy_config.cfproxy_worker_domains) + for dc, fallback_dst in cf_fallbacks.items(): + self._schedule_refill(dc, fallback_dst, worker_domains) log.info("CF worker pool warmup started for %d DC(s)", len(cf_fallbacks)) def reset(self): self._idle.clear() self._refilling.clear() + self._exhausted_until.clear() ws_pool = _WsPool()