17d211bf93
- Introduced a new web search block plugin to handle web search queries and results. - Updated copilot, doc block, and pro block plugins to include web search context in AI completions. - Implemented utility functions for parsing and building web search markdown. - Enhanced API to support web search requests and responses. - Added configuration for web search URL and timeout settings. - Updated size limit checks to account for web search content.
358 lines
14 KiB
Python
358 lines
14 KiB
Python
import asyncio
|
|
import hashlib
|
|
import math
|
|
import os
|
|
import time
|
|
from dataclasses import dataclass
|
|
from datetime import date
|
|
from typing import Any
|
|
|
|
from risk_config import RiskConfig
|
|
|
|
try: # pragma: no cover
|
|
from redis import asyncio as redis_asyncio
|
|
except Exception: # pragma: no cover
|
|
redis_asyncio = None
|
|
|
|
|
|
def _now_ms() -> int:
|
|
return int(time.time() * 1000)
|
|
|
|
|
|
def _utc_day() -> str:
|
|
return date.today().isoformat()
|
|
|
|
|
|
def stable_hash(value: str) -> str:
|
|
return hashlib.sha256(value.encode("utf-8")).hexdigest()
|
|
|
|
|
|
def estimate_tokens(text: str) -> int:
|
|
if not text:
|
|
return 0
|
|
ascii_chars = sum(1 for ch in text if ord(ch) < 128)
|
|
non_ascii = len(text) - ascii_chars
|
|
ascii_tokens = math.ceil(ascii_chars / 4)
|
|
non_ascii_tokens = math.ceil(non_ascii * 1.5)
|
|
return max(ascii_tokens + non_ascii_tokens, 1)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class RiskIdentity:
|
|
request_id: str
|
|
session_hash: str
|
|
ip_hash: str
|
|
route: str
|
|
method: str
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class RiskDecision:
|
|
allowed: bool
|
|
status_code: int = 200
|
|
reason: str = ""
|
|
error_code: str = ""
|
|
retry_after_seconds: int = 0
|
|
delay_ms: int = 0
|
|
|
|
|
|
class RiskRejected(RuntimeError):
|
|
def __init__(self, decision: RiskDecision) -> None:
|
|
super().__init__(decision.reason or decision.error_code or "request rejected")
|
|
self.decision = decision
|
|
|
|
|
|
class BaseRiskBackend:
|
|
async def incr_window(self, key: str, ttl_seconds: int) -> int:
|
|
raise NotImplementedError
|
|
|
|
async def get_float(self, key: str) -> float:
|
|
raise NotImplementedError
|
|
|
|
async def add_float(self, key: str, value: float, ttl_seconds: int) -> float:
|
|
raise NotImplementedError
|
|
|
|
async def get_int(self, key: str) -> int:
|
|
raise NotImplementedError
|
|
|
|
async def set_int(self, key: str, value: int, ttl_seconds: int) -> None:
|
|
raise NotImplementedError
|
|
|
|
async def set_float(self, key: str, value: float, ttl_seconds: int) -> None:
|
|
raise NotImplementedError
|
|
|
|
async def acquire_lock(self, key: str, ttl_seconds: int) -> bool:
|
|
raise NotImplementedError
|
|
|
|
async def release_lock(self, key: str) -> None:
|
|
raise NotImplementedError
|
|
|
|
|
|
class InMemoryRiskBackend(BaseRiskBackend):
|
|
def __init__(self) -> None:
|
|
self.values: dict[str, tuple[float, float]] = {}
|
|
self.locks: dict[str, float] = {}
|
|
self.guard = asyncio.Lock()
|
|
|
|
def _purge(self) -> None:
|
|
now = time.time()
|
|
for key, (_, expires_at) in list(self.values.items()):
|
|
if expires_at and expires_at <= now:
|
|
self.values.pop(key, None)
|
|
for key, expires_at in list(self.locks.items()):
|
|
if expires_at <= now:
|
|
self.locks.pop(key, None)
|
|
|
|
async def incr_window(self, key: str, ttl_seconds: int) -> int:
|
|
async with self.guard:
|
|
self._purge()
|
|
value, _ = self.values.get(key, (0.0, 0.0))
|
|
next_value = int(value) + 1
|
|
self.values[key] = (float(next_value), time.time() + ttl_seconds)
|
|
return next_value
|
|
|
|
async def get_float(self, key: str) -> float:
|
|
async with self.guard:
|
|
self._purge()
|
|
return float(self.values.get(key, (0.0, 0.0))[0])
|
|
|
|
async def add_float(self, key: str, value: float, ttl_seconds: int) -> float:
|
|
async with self.guard:
|
|
self._purge()
|
|
current, _ = self.values.get(key, (0.0, 0.0))
|
|
next_value = current + value
|
|
self.values[key] = (next_value, time.time() + ttl_seconds)
|
|
return next_value
|
|
|
|
async def get_int(self, key: str) -> int:
|
|
return int(await self.get_float(key))
|
|
|
|
async def set_int(self, key: str, value: int, ttl_seconds: int) -> None:
|
|
async with self.guard:
|
|
self._purge()
|
|
self.values[key] = (float(value), time.time() + ttl_seconds)
|
|
|
|
async def set_float(self, key: str, value: float, ttl_seconds: int) -> None:
|
|
async with self.guard:
|
|
self._purge()
|
|
self.values[key] = (float(value), time.time() + ttl_seconds)
|
|
|
|
async def acquire_lock(self, key: str, ttl_seconds: int) -> bool:
|
|
async with self.guard:
|
|
self._purge()
|
|
if key in self.locks:
|
|
return False
|
|
self.locks[key] = time.time() + ttl_seconds
|
|
return True
|
|
|
|
async def release_lock(self, key: str) -> None:
|
|
async with self.guard:
|
|
self.locks.pop(key, None)
|
|
|
|
|
|
class RedisRiskBackend(BaseRiskBackend):
|
|
def __init__(self, redis_url: str) -> None:
|
|
if redis_asyncio is None:
|
|
raise RuntimeError("redis package is not installed")
|
|
self.redis = redis_asyncio.from_url(redis_url, encoding="utf-8", decode_responses=True)
|
|
|
|
async def incr_window(self, key: str, ttl_seconds: int) -> int:
|
|
value = await self.redis.incr(key)
|
|
if value == 1:
|
|
await self.redis.expire(key, ttl_seconds)
|
|
return int(value)
|
|
|
|
async def get_float(self, key: str) -> float:
|
|
value = await self.redis.get(key)
|
|
if value is None:
|
|
return 0.0
|
|
return float(value)
|
|
|
|
async def add_float(self, key: str, value: float, ttl_seconds: int) -> float:
|
|
current = await self.get_float(key)
|
|
next_value = current + value
|
|
await self.redis.set(key, next_value, ex=ttl_seconds)
|
|
return next_value
|
|
|
|
async def get_int(self, key: str) -> int:
|
|
value = await self.redis.get(key)
|
|
if value is None:
|
|
return 0
|
|
try:
|
|
return int(value)
|
|
except (TypeError, ValueError):
|
|
return int(float(value))
|
|
|
|
async def set_int(self, key: str, value: int, ttl_seconds: int) -> None:
|
|
await self.redis.set(key, value, ex=ttl_seconds)
|
|
|
|
async def set_float(self, key: str, value: float, ttl_seconds: int) -> None:
|
|
await self.redis.set(key, value, ex=ttl_seconds)
|
|
|
|
async def acquire_lock(self, key: str, ttl_seconds: int) -> bool:
|
|
return bool(await self.redis.set(key, "1", ex=ttl_seconds, nx=True))
|
|
|
|
async def release_lock(self, key: str) -> None:
|
|
await self.redis.delete(key)
|
|
|
|
|
|
class RiskController:
|
|
def __init__(self, config: RiskConfig) -> None:
|
|
self.config = config
|
|
self.prefix = "llmtext:risk"
|
|
redis_url = os.getenv("REDIS_URL", "").strip()
|
|
if redis_url and redis_asyncio is not None:
|
|
self.backend: BaseRiskBackend = RedisRiskBackend(redis_url)
|
|
else:
|
|
self.backend = InMemoryRiskBackend()
|
|
|
|
def _api_key(self, identity: RiskIdentity, scope: str) -> str:
|
|
return f"{self.prefix}:api:{scope}:{identity.session_hash}:{identity.ip_hash}"
|
|
|
|
def _llm_key(self, identity: RiskIdentity, scope: str) -> str:
|
|
return f"{self.prefix}:llm:{scope}:{identity.session_hash}:{identity.ip_hash}"
|
|
|
|
def _budget_key(self, scope: str, scope_hash: str, current_day: str) -> str:
|
|
return f"{self.prefix}:budget:{scope}:{scope_hash}:{current_day}"
|
|
|
|
def _lock_key(self, scope: str, scope_hash: str) -> str:
|
|
return f"{self.prefix}:lock:{scope}:{scope_hash}"
|
|
|
|
def _circuit_key(self, scope: str) -> str:
|
|
return f"{self.prefix}:circuit:{scope}"
|
|
|
|
def _failure_key(self, model: str) -> str:
|
|
return f"{self.prefix}:failure:{model}"
|
|
|
|
async def check_api(self, identity: RiskIdentity, *, scope: str = "default") -> RiskDecision:
|
|
key = self._api_key(identity, scope)
|
|
count = await self.backend.incr_window(key, self.config.api_window_seconds)
|
|
if count > self.config.api_hard_limit_per_window:
|
|
return RiskDecision(
|
|
allowed=False,
|
|
status_code=429,
|
|
reason="请求过于频繁,请稍后再试",
|
|
error_code="api_rate_limited",
|
|
retry_after_seconds=self.config.api_window_seconds,
|
|
)
|
|
if count > self.config.api_soft_limit_per_window:
|
|
overflow = count - self.config.api_soft_limit_per_window
|
|
delay_ms = min(self.config.delay_cap_ms, overflow * self.config.delay_step_ms)
|
|
return RiskDecision(allowed=True, delay_ms=delay_ms)
|
|
return RiskDecision(allowed=True)
|
|
|
|
async def check_llm(
|
|
self,
|
|
identity: RiskIdentity,
|
|
*,
|
|
scope: str,
|
|
estimated_cost: float,
|
|
) -> RiskDecision:
|
|
global_circuit = await self.backend.get_int(self._circuit_key("global"))
|
|
model_circuit = await self.backend.get_int(self._circuit_key(scope))
|
|
if global_circuit > 0 or model_circuit > 0:
|
|
return RiskDecision(
|
|
allowed=False,
|
|
status_code=503,
|
|
reason="当前推理服务繁忙,请稍后再试",
|
|
error_code="llm_circuit_open",
|
|
retry_after_seconds=self.config.model_circuit_ttl_seconds,
|
|
)
|
|
if estimated_cost > self.config.single_request_max_cost_usd:
|
|
return RiskDecision(
|
|
allowed=False,
|
|
status_code=429,
|
|
reason="单次请求成本过高,已被拒绝",
|
|
error_code="llm_cost_too_high",
|
|
)
|
|
day = _utc_day()
|
|
session_budget = await self.backend.get_float(self._budget_key("session", identity.session_hash, day))
|
|
ip_budget = await self.backend.get_float(self._budget_key("ip", identity.ip_hash, day))
|
|
global_budget = await self.backend.get_float(self._budget_key("global", "global", day))
|
|
if session_budget + estimated_cost > self.config.daily_budget_session_usd:
|
|
return RiskDecision(False, 429, "当前匿名会话今日额度已用尽", "session_budget_exhausted", 3600)
|
|
if ip_budget + estimated_cost > self.config.daily_budget_ip_usd:
|
|
return RiskDecision(False, 429, "当前网络环境今日额度已用尽", "ip_budget_exhausted", 3600)
|
|
if global_budget + estimated_cost > self.config.daily_budget_global_usd:
|
|
await self.backend.set_int(self._circuit_key("global"), 1, self.config.model_circuit_ttl_seconds)
|
|
return RiskDecision(False, 503, "今日全局推理预算已耗尽", "global_budget_exhausted", 3600)
|
|
key = self._llm_key(identity, scope)
|
|
count = await self.backend.incr_window(key, self.config.llm_window_seconds)
|
|
if count > self.config.llm_hard_limit_per_window:
|
|
return RiskDecision(False, 429, "推理请求过于频繁,请稍后重试", "llm_rate_limited", self.config.llm_window_seconds)
|
|
if count > self.config.llm_soft_limit_per_window:
|
|
overflow = count - self.config.llm_soft_limit_per_window
|
|
delay_ms = min(self.config.delay_cap_ms, overflow * self.config.delay_step_ms)
|
|
return RiskDecision(True, delay_ms=delay_ms)
|
|
return RiskDecision(True)
|
|
|
|
async def reserve_budget(self, identity: RiskIdentity, estimated_cost: float) -> None:
|
|
day = _utc_day()
|
|
ttl_seconds = 60 * 60 * 24
|
|
await self.backend.add_float(self._budget_key("session", identity.session_hash, day), estimated_cost, ttl_seconds)
|
|
await self.backend.add_float(self._budget_key("ip", identity.ip_hash, day), estimated_cost, ttl_seconds)
|
|
await self.backend.add_float(self._budget_key("global", "global", day), estimated_cost, ttl_seconds)
|
|
|
|
async def acquire_execution_slot(self, identity: RiskIdentity, *, model: str) -> list[str]:
|
|
ttl_seconds = 60 * 15
|
|
keys = [
|
|
self._lock_key("session", f"{identity.session_hash}:{identity.request_id}"),
|
|
self._lock_key("global", identity.request_id),
|
|
]
|
|
session_running = await self.backend.get_int(self._lock_key("session-count", identity.session_hash))
|
|
global_running = await self.backend.get_int(self._lock_key("global-count", "global"))
|
|
if session_running >= self.config.session_concurrency_limit:
|
|
raise RiskRejected(
|
|
RiskDecision(False, 429, "当前会话并发推理过多,请稍后重试", "session_concurrency_limited", 30)
|
|
)
|
|
if global_running >= self.config.global_concurrency_limit:
|
|
raise RiskRejected(
|
|
RiskDecision(False, 503, "当前全局推理负载过高,请稍后重试", "global_concurrency_limited", 30)
|
|
)
|
|
acquired: list[str] = []
|
|
for key in keys:
|
|
ok = await self.backend.acquire_lock(key, ttl_seconds)
|
|
if not ok:
|
|
for acquired_key in acquired:
|
|
await self.backend.release_lock(acquired_key)
|
|
raise RiskRejected(
|
|
RiskDecision(False, 429, "当前请求正在执行,请勿重复提交", "duplicate_request", 10)
|
|
)
|
|
acquired.append(key)
|
|
await self.backend.add_float(self._lock_key("session-count", identity.session_hash), 1.0, ttl_seconds)
|
|
await self.backend.add_float(self._lock_key("global-count", "global"), 1.0, ttl_seconds)
|
|
return acquired
|
|
|
|
async def release_execution_slot(self, identity: RiskIdentity, lock_keys: list[str], *, model: str) -> None:
|
|
for key in lock_keys:
|
|
await self.backend.release_lock(key)
|
|
session_count_key = self._lock_key("session-count", identity.session_hash)
|
|
global_count_key = self._lock_key("global-count", "global")
|
|
session_count = max(0.0, await self.backend.get_float(session_count_key) - 1.0)
|
|
global_count = max(0.0, await self.backend.get_float(global_count_key) - 1.0)
|
|
await self.backend.set_float(session_count_key, session_count, 60 * 15)
|
|
await self.backend.set_float(global_count_key, global_count, 60 * 15)
|
|
|
|
async def record_model_result(self, *, model: str, success: bool) -> None:
|
|
if success:
|
|
await self.backend.set_int(self._failure_key(model), 0, self.config.model_circuit_ttl_seconds)
|
|
return
|
|
failures = await self.backend.incr_window(self._failure_key(model), self.config.model_circuit_ttl_seconds)
|
|
if failures >= self.config.model_circuit_breaker_failures:
|
|
await self.backend.set_int(self._circuit_key(model), 1, self.config.model_circuit_ttl_seconds)
|
|
|
|
|
|
_risk_controller: RiskController | None = None
|
|
|
|
|
|
def get_risk_controller(config: RiskConfig) -> RiskController:
|
|
global _risk_controller
|
|
if _risk_controller is None:
|
|
_risk_controller = RiskController(config)
|
|
return _risk_controller
|
|
|
|
|
|
def reset_risk_controller() -> None:
|
|
global _risk_controller
|
|
_risk_controller = None
|