deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/runtime/runs/concurrency_gate.py
2026-09-07 18:24:55 +08:00

380 lines
13 KiB
Python

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