CF Worker pool refactoring

This commit is contained in:
Flowseal
2026-07-31 16:38:42 +03:00
parent 41e97c6a05
commit aee473c9f9
2 changed files with 98 additions and 61 deletions
+24 -21
View File
@@ -1,7 +1,6 @@
import asyncio import asyncio
import logging import logging
import struct import struct
import random
from typing import List, Optional from typing import List, Optional
from urllib.parse import urlencode 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 worker_domains = proxy_config.cfproxy_worker_domains
if not worker_domains: if not worker_domains:
return False return False
random.shuffle(worker_domains)
for worker_domain in worker_domains: pooled = None if is_test_dc else await cf_worker_pool.get(
ws = None if is_test_dc else await cf_worker_pool.get(dc, worker_domain, fallback_dst) dc, fallback_dst, worker_domains)
if ws: if pooled:
log.info("[%s] DC%d%s -> CF worker pool hit for %s", ws, worker_domain = pooled
label, dc, media_tag, fallback_dst) log.info("[%s] DC%d%s -> CF worker pool hit via %s for %s",
else: label, dc, media_tag, worker_domain, fallback_dst)
query = urlencode({ else:
'dst': fallback_dst, query = urlencode({
'dc': str(dc), 'dst': fallback_dst,
}) 'dc': str(dc),
path = f'/apiws?{query}' })
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", log.info("[%s] DC%d%s -> trying CF worker %s for %s",
label, dc, media_tag, worker_domain, fallback_dst) 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, ws = await RawWebSocket.connect(worker_domain, worker_domain,
timeout=10.0, path=path) timeout=10.0, path=path)
except Exception as exc: except Exception as exc:
cf_worker_pool.report_failure(worker_domain, exc)
log.warning("[%s] DC%d%s CF worker %s failed: %s", log.warning("[%s] DC%d%s CF worker %s failed: %s",
label, dc, media_tag, worker_domain, repr(exc)) label, dc, media_tag, worker_domain, repr(exc))
continue continue
stats.connections_cfproxy += 1 if ws is None:
await ws.send(relay_init) return False
await bridge_ws_reencrypt(reader, writer, ws, label, ctx,
dc=dc, is_media=is_media, stats.connections_cfproxy += 1
splitter=None) await ws.send(relay_init)
return True await bridge_ws_reencrypt(reader, writer, ws, label, ctx,
return False dc=dc, is_media=is_media,
splitter=None)
return True
async def _cfproxy_fallback(reader, writer, relay_init, label, async def _cfproxy_fallback(reader, writer, relay_init, label,
+74 -40
View File
@@ -1,5 +1,6 @@
import asyncio import asyncio
import logging import logging
import random
import time import time
from collections import deque from collections import deque
@@ -221,73 +222,105 @@ class _CfWorkerPool:
PER_DC_LIMIT = 1 PER_DC_LIMIT = 1
def __init__(self): def __init__(self):
self._idle: Dict[Tuple[int, str], deque] = {} self._idle: Dict[int, deque] = {}
self._refilling: Set[Tuple[int, str]] = set() 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() now = time.monotonic()
key = (dc, worker_domain)
bucket = self._idle.get(key) bucket = self._idle.get(dc)
if bucket is None: if bucket is None:
bucket = deque() bucket = deque()
self._idle[key] = bucket self._idle[dc] = bucket
while bucket: while bucket:
ws, created = bucket.popleft() ws, created, worker_domain = bucket.popleft()
age = now - created age = now - created
if (age > self.WS_POOL_MAX_AGE or ws._closed if (age > self.WS_POOL_MAX_AGE or ws._closed
or ws.writer.transport.is_closing()): or ws.writer.transport.is_closing()):
asyncio.create_task(self._quiet_close(ws)) asyncio.create_task(self._quiet_close(ws))
continue continue
stats.cf_pool_hits += 1 stats.cf_pool_hits += 1
log.debug("CF worker pool hit DC%d (age=%.1fs, left=%d)", log.debug(
dc, age, len(bucket)) "CF worker pool hit DC%d via %s (age=%.1fs, left=%d)",
self._schedule_refill(key, fallback_dst) dc, worker_domain, age, len(bucket))
return ws self._schedule_refill(dc, fallback_dst, worker_domains)
return ws, worker_domain
stats.cf_pool_misses += 1 stats.cf_pool_misses += 1
self._schedule_refill(key, fallback_dst)
return None return None
def _schedule_refill(self, key, fallback_dst): def _schedule_refill(self, dc, fallback_dst, worker_domains):
if key in self._refilling: if dc in self._refilling:
return return
self._refilling.add(key) self._refilling.add(dc)
asyncio.create_task(self._refill(key, fallback_dst)) asyncio.create_task(self._refill(
dc, fallback_dst, list(worker_domains)))
async def _refill(self, key, fallback_dst): async def _refill(self, dc, fallback_dst, worker_domains):
dc, worker_domain = key
try: try:
bucket = self._idle.setdefault(key, deque()) bucket = self._idle.setdefault(dc, deque())
needed = min(proxy_config.pool_size - len(bucket), self.PER_DC_LIMIT) target_size = min(proxy_config.pool_size, self.PER_DC_LIMIT)
needed = target_size - len(bucket)
if needed <= 0: if needed <= 0:
return return
tasks = [asyncio.create_task(
self._connect_one(worker_domain, fallback_dst, dc)) for _ in range(needed):
for _ in range(needed)] connected = await self._connect_one(
for t in tasks: worker_domains, fallback_dst, dc)
try: if connected is None:
ws = await t break
if ws: ws, worker_domain = connected
bucket.append((ws, time.monotonic())) bucket.append((ws, time.monotonic(), worker_domain))
except Exception:
pass
log.debug("CF worker pool refilled DC%d: %d ready", log.debug("CF worker pool refilled DC%d: %d ready",
dc, len(bucket)) dc, len(bucket))
finally: 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({ query = urlencode({
'dst': fallback_dst, 'dst': fallback_dst,
'dc': str(dc), 'dc': str(dc),
}) })
path = f'/apiws?{query}' path = f'/apiws?{query}'
try: for worker_domain in self.available_domains(worker_domains):
return await RawWebSocket.connect( try:
worker_domain, worker_domain, timeout=8, path=path) ws = await RawWebSocket.connect(
except Exception: worker_domain, worker_domain, timeout=8, path=path)
return None 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): async def _quiet_close(self, ws):
try: try:
@@ -304,15 +337,16 @@ class _CfWorkerPool:
if not cf_fallbacks or not proxy_config.cfproxy_worker_domains: if not cf_fallbacks or not proxy_config.cfproxy_worker_domains:
return return
for worker_domain in proxy_config.cfproxy_worker_domains: worker_domains = list(proxy_config.cfproxy_worker_domains)
for dc, fallback_dst in cf_fallbacks.items(): for dc, fallback_dst in cf_fallbacks.items():
self._schedule_refill((dc, worker_domain), fallback_dst) self._schedule_refill(dc, fallback_dst, worker_domains)
log.info("CF worker pool warmup started for %d DC(s)", len(cf_fallbacks)) log.info("CF worker pool warmup started for %d DC(s)", len(cf_fallbacks))
def reset(self): def reset(self):
self._idle.clear() self._idle.clear()
self._refilling.clear() self._refilling.clear()
self._exhausted_until.clear()
ws_pool = _WsPool() ws_pool = _WsPool()