refactor: unify LLM gateway policy

This commit is contained in:
leefer
2026-07-29 20:08:39 +08:00
parent 5f0b9735bd
commit a368174a75
10 changed files with 587 additions and 224 deletions
+15
View File
@@ -0,0 +1,15 @@
from .gateway import (
LLMGateway,
LLMGatewayError,
LLMResult,
LLMStreamEvent,
ModelProfile,
)
__all__ = [
"LLMGateway",
"LLMGatewayError",
"LLMResult",
"LLMStreamEvent",
"ModelProfile",
]
+254
View File
@@ -0,0 +1,254 @@
from __future__ import annotations
import time
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Generic, TypeVar
T = TypeVar("T")
class LLMGatewayError(ValueError):
"""Stable application error that does not expose provider details."""
def __init__(self, message: str, code: str = "unavailable") -> None:
super().__init__(message)
self.code = code
@dataclass(frozen=True)
class ModelProfile:
role: str
api_key: str
base_url: str
model: str
@property
def configured(self) -> bool:
return bool(self.api_key and self.base_url and self.model)
@dataclass(frozen=True)
class LLMResult(Generic[T]):
value: T
source: str
role: str
model: str
latency_ms: int
@dataclass(frozen=True)
class LLMStreamEvent(Generic[T]):
kind: str
value: T | None = None
source: str = ""
role: str = ""
model: str = ""
latency_ms: int = 0
class LLMGateway:
"""Single policy boundary for access, model fallback, and call auditing."""
def __init__(
self,
*,
database: Any,
user_id_supplier: Callable[[], int],
membership_supplier: Callable[[], dict[str, Any]],
settings_supplier: Callable[[], dict[str, Any]],
profile_supplier: Callable[[], dict[str, Any]],
) -> None:
self.database = database
self.user_id_supplier = user_id_supplier
self.membership_supplier = membership_supplier
self.settings_supplier = settings_supplier
self.profile_supplier = profile_supplier
def ensure_access(self, feature: str) -> tuple[str, tuple[ModelProfile, ...]]:
del feature # Reserved for future feature-specific policy.
membership = self.membership_supplier()
profile = self.profile_supplier()
source = str(profile.get("source") or "none")
profiles = self._model_profiles(profile)
if source == "none" or not profiles:
raise LLMGatewayError("智能功能尚未配置,请联系管理员。", "not_configured")
if source == "platform":
settings = self.settings_supplier()
limit = max(1, int(settings.get("member_daily_limit") or 50))
if not membership.get("active"):
raise LLMGatewayError("开通会员后可使用智能功能。", "membership_required")
if self._usage_today(source) >= limit:
raise LLMGatewayError(
f"今日会员模型额度已用完({limit} 次)。", "quota_exhausted"
)
return source, profiles
def call(
self,
feature: str,
prompt_version: str,
invoke: Callable[[ModelProfile], T],
error_types: tuple[type[BaseException], ...],
) -> LLMResult[T]:
source, profiles = self.ensure_access(feature)
started = time.perf_counter()
last_error: BaseException | None = None
for profile in profiles:
try:
value = invoke(profile)
except error_types as exc:
last_error = exc
continue
latency_ms = round((time.perf_counter() - started) * 1000)
self.audit(
feature, source, profile, "success", latency_ms, prompt_version
)
return LLMResult(
value=value,
source=source,
role=profile.role,
model=profile.model,
latency_ms=latency_ms,
)
failed = profiles[-1]
latency_ms = round((time.perf_counter() - started) * 1000)
self.audit(
feature,
source,
failed,
"failed",
latency_ms,
prompt_version,
self._error_code(last_error),
)
raise LLMGatewayError("智能解读服务暂不可用,请稍后重试。") from last_error
@staticmethod
def probe(profile: dict[str, Any], invoke: Callable[[ModelProfile], T]) -> T:
"""Route an explicit administrator connection test through the gateway boundary."""
model = ModelProfile(
role="probe",
api_key=str(profile.get("api_key") or ""),
base_url=str(profile.get("base_url") or ""),
model=str(profile.get("model") or ""),
)
return invoke(model)
def stream(
self,
feature: str,
prompt_version: str,
invoke: Callable[[ModelProfile], Iterator[T]],
error_types: tuple[type[BaseException], ...],
) -> Iterator[LLMStreamEvent[T]]:
source, profiles = self.ensure_access(feature)
started = time.perf_counter()
last_error: BaseException | None = None
for profile in profiles:
try:
upstream = iter(invoke(profile))
first = next(upstream)
except (*error_types, StopIteration) as exc:
last_error = exc
continue
yield LLMStreamEvent(kind="delta", value=first)
try:
for chunk in upstream:
yield LLMStreamEvent(kind="delta", value=chunk)
except error_types as exc:
latency_ms = round((time.perf_counter() - started) * 1000)
self.audit(
feature,
source,
profile,
"failed",
latency_ms,
prompt_version,
self._error_code(exc),
)
raise LLMGatewayError(
"智能解读连接中断,请稍后重试。"
) from exc
latency_ms = round((time.perf_counter() - started) * 1000)
self.audit(
feature, source, profile, "success", latency_ms, prompt_version
)
yield LLMStreamEvent(
kind="complete",
source=source,
role=profile.role,
model=profile.model,
latency_ms=latency_ms,
)
return
failed = profiles[-1]
latency_ms = round((time.perf_counter() - started) * 1000)
self.audit(
feature,
source,
failed,
"failed",
latency_ms,
prompt_version,
self._error_code(last_error),
)
raise LLMGatewayError("智能解读服务暂不可用,请稍后重试。") from last_error
def audit(
self,
feature: str,
source: str,
profile: ModelProfile,
status: str,
latency_ms: int,
prompt_version: str,
error_code: str = "",
input_tokens: int = 0,
output_tokens: int = 0,
) -> None:
self.database.record_llm_usage(
self.user_id_supplier(),
feature,
source,
profile.model,
status,
latency_ms,
role=profile.role,
prompt_version=prompt_version,
error_code=error_code,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
def _usage_today(self, source: str) -> int:
now = datetime.now().astimezone()
start = now.replace(
hour=0, minute=0, second=0, microsecond=0
).astimezone(timezone.utc)
return self.database.count_llm_usage_since(
self.user_id_supplier(), source, start.isoformat(timespec="seconds")
)
@staticmethod
def _model_profiles(profile: dict[str, Any]) -> tuple[ModelProfile, ...]:
result = []
for role in ("primary", "fallback"):
item = profile.get(role) or {}
candidate = ModelProfile(
role=role,
api_key=str(item.get("api_key") or ""),
base_url=str(item.get("base_url") or ""),
model=str(item.get("model") or ""),
)
if candidate.configured:
result.append(candidate)
return tuple(result)
@staticmethod
def _error_code(error: BaseException | None) -> str:
if error is None:
return "empty_response"
return type(error).__name__[:80]