"""Optional cross-worker run concurrency gates. The Redis gate is deliberately fail-open by default: Redis coordinates global limits when it is fast and healthy, but it must not block the chat path when it is slow or unavailable. """ from __future__ import annotations import asyncio import logging import time from dataclasses import dataclass from enum import Enum from typing import Any from deerflow.config.concurrency_config import RedisConcurrencyConfig logger = logging.getLogger(__name__) class AcquireDecision(str, Enum): ACQUIRED = "acquired" REJECTED = "rejected" DEGRADED = "degraded" @dataclass(frozen=True) class AcquireResult: decision: AcquireDecision reason: str = "" source: str = "redis" @property def acquired(self) -> bool: return self.decision == AcquireDecision.ACQUIRED @property def rejected(self) -> bool: return self.decision == AcquireDecision.REJECTED @property def degraded(self) -> bool: return self.decision == AcquireDecision.DEGRADED @dataclass(frozen=True) class ActiveLease: run_id: str thread_id: str = "" model_name: str = "" user_id: str = "" assistant_id: str = "" created_at: float = 0.0 expires_at: float = 0.0 class CircuitBreaker: """Tiny in-process circuit breaker for Redis failures.""" def __init__(self, *, failure_threshold: int, cooldown_seconds: int) -> None: self._failure_threshold = max(1, int(failure_threshold)) self._cooldown_seconds = max(1, int(cooldown_seconds)) self._failures = 0 self._opened_until = 0.0 def allow(self) -> bool: return time.monotonic() >= self._opened_until def record_success(self) -> None: self._failures = 0 self._opened_until = 0.0 def record_failure(self) -> None: self._failures += 1 if self._failures >= self._failure_threshold: self._opened_until = time.monotonic() + self._cooldown_seconds @property def open(self) -> bool: return not self.allow() class RedisConcurrencyGate: """Redis lease-based global concurrency limiter.""" _ACQUIRE_SCRIPT = """ local now = tonumber(ARGV[1]) local expires_at = tonumber(ARGV[2]) local run_id = ARGV[3] local model_limit = tonumber(ARGV[4]) local user_model_limit = tonumber(ARGV[5]) local reject_thread = tonumber(ARGV[6]) local payload = ARGV[7] local ttl_ms = tonumber(ARGV[8]) for i = 1, 4 do redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now) end if reject_thread == 1 and redis.call('ZCARD', KEYS[4]) > 0 then return {'rejected', 'thread'} end if model_limit > 0 and redis.call('ZCARD', KEYS[2]) >= model_limit then return {'rejected', 'model'} end if user_model_limit > 0 and redis.call('ZCARD', KEYS[3]) >= user_model_limit then return {'rejected', 'user_model'} end redis.call('ZADD', KEYS[1], expires_at, run_id) redis.call('ZADD', KEYS[2], expires_at, run_id) if user_model_limit > 0 then redis.call('ZADD', KEYS[3], expires_at, run_id) end redis.call('ZADD', KEYS[4], expires_at, run_id) redis.call('SET', KEYS[5], payload, 'PX', ttl_ms) return {'acquired', ''} """ def __init__(self, config: RedisConcurrencyConfig) -> None: self.config = config self._client: Any | None = None self._breaker = CircuitBreaker( failure_threshold=config.circuit_breaker_failures, cooldown_seconds=config.circuit_breaker_cooldown_seconds, ) @property def enabled(self) -> bool: return bool(self.config.enabled and self.config.url) def health(self) -> dict[str, Any]: return { "enabled": self.enabled, "source": "redis", "circuit_open": self._breaker.open, "fail_open": self.config.fail_open, } async def aclose(self) -> None: client = self._client self._client = None if client is not None: try: await client.aclose() except Exception: logger.debug("Failed to close Redis concurrency client", exc_info=True) def _ns(self) -> str: return self.config.namespace.strip().rstrip(":") or "cmzs_deerflow:concurrency" def _keys(self, *, run_id: str, model_name: str, user_id: str, thread_id: str) -> list[str]: ns = self._ns() return [ f"{ns}:all", f"{ns}:model:{model_name}", f"{ns}:user_model:{model_name}:{user_id}", f"{ns}:thread:{thread_id}", f"{ns}:run:{run_id}", ] async def _redis(self) -> Any: if self._client is not None: return self._client try: from redis.asyncio import Redis except Exception as exc: # pragma: no cover - exercised when dependency is absent raise RuntimeError("redis package is not installed") from exc timeout = max(0.001, self.config.socket_timeout_ms / 1000) connect_timeout = max(0.001, self.config.socket_connect_timeout_ms / 1000) self._client = Redis.from_url( self.config.url, password=self.config.password, socket_timeout=timeout, socket_connect_timeout=connect_timeout, decode_responses=True, ) return self._client async def _with_timeout(self, coro, timeout_ms: int) -> Any: return await asyncio.wait_for(coro, timeout=max(0.001, timeout_ms / 1000)) def _degraded(self, reason: str) -> AcquireResult: return AcquireResult(AcquireDecision.DEGRADED, reason=reason) def _on_failure(self, action: str, exc: BaseException) -> None: self._breaker.record_failure() logger.warning( "Redis concurrency %s failed; fail_open=%s circuit_open=%s: %s", action, self.config.fail_open, self._breaker.open, exc, ) async def try_acquire( self, *, run_id: str, thread_id: str, model_name: str, user_id: str, assistant_id: str = "", model_limit: int | None = None, user_model_limit: int | None = None, reject_thread: bool = True, ) -> AcquireResult: if not self.enabled: return self._degraded("disabled") if not self._breaker.allow(): return self._degraded("circuit_open") now = time.time() expires_at = now + self.config.lease_ttl_seconds payload = { "run_id": run_id, "thread_id": thread_id, "model_name": model_name, "user_id": user_id, "assistant_id": assistant_id, "created_at": now, "expires_at": expires_at, } import json keys = self._keys(run_id=run_id, model_name=model_name, user_id=user_id, thread_id=thread_id) ttl_ms = int(self.config.lease_ttl_seconds * 1000) try: client = await self._redis() raw = await self._with_timeout( client.eval( self._ACQUIRE_SCRIPT, len(keys), *keys, str(now), str(expires_at), run_id, str(max(0, int(model_limit or 0))), str(max(0, int(user_model_limit or 0))), "1" if reject_thread else "0", json.dumps(payload, ensure_ascii=False), str(ttl_ms), ), self.config.acquire_timeout_ms, ) self._breaker.record_success() except Exception as exc: self._on_failure("acquire", exc) if self.config.fail_open: return self._degraded(type(exc).__name__) raise status = raw[0] if isinstance(raw, list) and raw else "" reason = raw[1] if isinstance(raw, list) and len(raw) > 1 else "" if status == "acquired": return AcquireResult(AcquireDecision.ACQUIRED) return AcquireResult(AcquireDecision.REJECTED, reason=str(reason or "limit")) async def heartbeat( self, *, run_id: str, thread_id: str, model_name: str, user_id: str, user_model_accounted: bool = True, ) -> None: if not self.enabled or not self._breaker.allow(): return now = time.time() expires_at = now + self.config.lease_ttl_seconds keys = self._keys(run_id=run_id, model_name=model_name, user_id=user_id, thread_id=thread_id) try: client = await self._redis() pipe = client.pipeline(transaction=False) refresh_keys = [keys[0], keys[1], keys[3]] if user_model_accounted: refresh_keys.append(keys[2]) for key in refresh_keys: pipe.zadd(key, {run_id: expires_at}) pipe.expire(keys[4], self.config.lease_ttl_seconds) await self._with_timeout(pipe.execute(), self.config.operation_timeout_ms) self._breaker.record_success() except Exception as exc: self._on_failure("heartbeat", exc) async def release( self, *, run_id: str, thread_id: str, model_name: str, user_id: str, user_model_accounted: bool = True, ) -> None: if not self.enabled or not self._breaker.allow(): return keys = self._keys(run_id=run_id, model_name=model_name, user_id=user_id, thread_id=thread_id) try: client = await self._redis() pipe = client.pipeline(transaction=False) release_keys = [keys[0], keys[1], keys[3]] if user_model_accounted: release_keys.append(keys[2]) for key in release_keys: pipe.zrem(key, run_id) pipe.delete(keys[4]) await self._with_timeout(pipe.execute(), self.config.operation_timeout_ms) self._breaker.record_success() except Exception as exc: self._on_failure("release", exc) async def count_active(self) -> int | None: if not self.enabled or not self._breaker.allow(): return None key = f"{self._ns()}:all" now = time.time() try: client = await self._redis() pipe = client.pipeline(transaction=False) pipe.zremrangebyscore(key, "-inf", now) pipe.zcard(key) result = await self._with_timeout(pipe.execute(), self.config.operation_timeout_ms) self._breaker.record_success() return int(result[-1] or 0) except Exception as exc: self._on_failure("count", exc) return None async def list_active(self, *, limit: int = 1000) -> list[ActiveLease] | None: if not self.enabled or not self._breaker.allow(): return None import json all_key = f"{self._ns()}:all" now = time.time() try: client = await self._redis() await self._with_timeout(client.zremrangebyscore(all_key, "-inf", now), self.config.operation_timeout_ms) run_ids = await self._with_timeout(client.zrange(all_key, 0, max(0, limit - 1)), self.config.operation_timeout_ms) if not run_ids: self._breaker.record_success() return [] keys = [f"{self._ns()}:run:{run_id}" for run_id in run_ids] payloads = await self._with_timeout(client.mget(keys), self.config.operation_timeout_ms) leases: list[ActiveLease] = [] for run_id, payload in zip(run_ids, payloads, strict=False): if not payload: continue try: data = json.loads(payload) except Exception: continue leases.append( ActiveLease( run_id=str(data.get("run_id") or run_id), thread_id=str(data.get("thread_id") or ""), model_name=str(data.get("model_name") or ""), user_id=str(data.get("user_id") or ""), assistant_id=str(data.get("assistant_id") or ""), created_at=float(data.get("created_at") or 0), expires_at=float(data.get("expires_at") or 0), ) ) self._breaker.record_success() return leases except Exception as exc: self._on_failure("list", exc) return None def make_redis_concurrency_gate(config: Any) -> RedisConcurrencyGate | None: concurrency = getattr(config, "concurrency", None) redis_config = getattr(concurrency, "redis", None) if redis_config is None or not getattr(redis_config, "enabled", False) or not getattr(redis_config, "url", None): return None return RedisConcurrencyGate(redis_config)