380 lines
13 KiB
Python
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)
|