rebuild(stage-10): deliver mentor and unified llm streaming
This commit is contained in:
@@ -0,0 +1,262 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from backend.database.connection import Database
|
||||
from backend.errors import BusinessError
|
||||
from backend.features.accounts.model_pool import ModelPoolService
|
||||
from backend.features.accounts.models import Principal
|
||||
from backend.features.accounts.service import MembershipService
|
||||
from backend.llm.provider import OpenAICompatibleClient, ProviderFailure
|
||||
from backend.llm.repository import LLMRepository
|
||||
|
||||
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
|
||||
|
||||
class LLMGatewayError(RuntimeError):
|
||||
def __init__(self, code: str, message: str, *, partial: bool = False) -> None:
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.partial = partial
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LLMProfile:
|
||||
model_id: int
|
||||
role: str
|
||||
base_url: str
|
||||
model_identifier: str
|
||||
api_key: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LLMCall:
|
||||
request_id: str
|
||||
user_id: int
|
||||
feature: str
|
||||
prompt_version: str
|
||||
business_id: str
|
||||
started_at: datetime
|
||||
usage_date: str
|
||||
input_chars: int
|
||||
quota_exempt: bool
|
||||
profiles: tuple[LLMProfile, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LLMStreamEvent:
|
||||
type: str
|
||||
content: str = ""
|
||||
request_id: str = ""
|
||||
|
||||
|
||||
class LLMGateway:
|
||||
def __init__(
|
||||
self,
|
||||
database: Database,
|
||||
repository: LLMRepository,
|
||||
memberships: MembershipService,
|
||||
model_pool: ModelPoolService,
|
||||
provider: OpenAICompatibleClient | None = None,
|
||||
) -> None:
|
||||
self._database = database
|
||||
self._repository = repository
|
||||
self._memberships = memberships
|
||||
self._model_pool = model_pool
|
||||
self._provider = provider or OpenAICompatibleClient()
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
principal: Principal,
|
||||
*,
|
||||
feature: str,
|
||||
prompt_version: str,
|
||||
business_id: str,
|
||||
input_chars: int,
|
||||
) -> LLMCall:
|
||||
membership = self._memberships.view_for(principal)
|
||||
if not self._memberships.can_use_smart_features(principal):
|
||||
raise LLMGatewayError("membership_required", "该功能仅对会员开放。")
|
||||
try:
|
||||
runtime = self._model_pool.runtime_config()
|
||||
except BusinessError as exc:
|
||||
raise LLMGatewayError("not_configured", "智能解读服务尚未配置。") from exc
|
||||
records = (("primary", runtime.primary), ("fallback", runtime.fallback))
|
||||
profiles = tuple(
|
||||
LLMProfile(
|
||||
model_id=record.id,
|
||||
role=role,
|
||||
base_url=record.base_url,
|
||||
model_identifier=record.model_identifier,
|
||||
api_key=self._model_pool.decrypt_api_key(record),
|
||||
)
|
||||
for role, record in records
|
||||
if record is not None
|
||||
)
|
||||
now = datetime.now(SHANGHAI)
|
||||
request_id = uuid.uuid4().hex
|
||||
with self._database.transaction() as connection:
|
||||
used = self._repository.usage_today(
|
||||
connection, principal.user.id, now.date().isoformat()
|
||||
)
|
||||
active = self._repository.active_today(
|
||||
connection, principal.user.id, now.date().isoformat()
|
||||
)
|
||||
if not membership.quota_exempt and used + active >= membership.daily_limit:
|
||||
raise LLMGatewayError(
|
||||
"quota_exhausted",
|
||||
f"今日智能分析额度已用完({membership.daily_limit} 次)。",
|
||||
)
|
||||
self._repository.reserve(
|
||||
connection,
|
||||
request_id=request_id,
|
||||
user_id=principal.user.id,
|
||||
feature=feature,
|
||||
business_id=business_id,
|
||||
prompt_version=prompt_version,
|
||||
started_at=now.isoformat(timespec="seconds"),
|
||||
input_chars=max(0, input_chars),
|
||||
)
|
||||
return LLMCall(
|
||||
request_id=request_id,
|
||||
user_id=principal.user.id,
|
||||
feature=feature,
|
||||
prompt_version=prompt_version,
|
||||
business_id=business_id,
|
||||
started_at=now,
|
||||
usage_date=now.date().isoformat(),
|
||||
input_chars=max(0, input_chars),
|
||||
quota_exempt=membership.quota_exempt,
|
||||
profiles=profiles,
|
||||
)
|
||||
|
||||
def stream(
|
||||
self, call: LLMCall, messages: list[dict[str, str]]
|
||||
) -> Iterator[LLMStreamEvent]:
|
||||
last_error = "unavailable"
|
||||
for profile in call.profiles:
|
||||
attempt_started = time.perf_counter()
|
||||
with self._database.transaction() as connection:
|
||||
attempt_id = self._repository.start_attempt(
|
||||
connection,
|
||||
call.request_id,
|
||||
profile.model_id,
|
||||
profile.role,
|
||||
_now(),
|
||||
call.input_chars,
|
||||
)
|
||||
output = ""
|
||||
try:
|
||||
upstream = iter(self._provider.stream(profile, messages))
|
||||
first = next(upstream)
|
||||
except (ProviderFailure, StopIteration) as exc:
|
||||
last_error = exc.code if isinstance(exc, ProviderFailure) else "empty_response"
|
||||
self._finish_attempt(attempt_id, "failed", attempt_started, last_error, 0)
|
||||
continue
|
||||
except Exception as exc:
|
||||
last_error = "unavailable"
|
||||
self._finish_attempt(attempt_id, "failed", attempt_started, last_error, 0)
|
||||
self._finish_request(call, "failed", last_error, 0)
|
||||
raise LLMGatewayError(
|
||||
"model_unavailable", "智能解读服务暂不可用,请稍后重试。"
|
||||
) from exc
|
||||
with self._database.transaction() as connection:
|
||||
self._repository.set_request_status(connection, call.request_id, "streaming")
|
||||
output = first
|
||||
try:
|
||||
yield LLMStreamEvent("delta", first, call.request_id)
|
||||
for chunk in upstream:
|
||||
output += chunk
|
||||
yield LLMStreamEvent("delta", chunk, call.request_id)
|
||||
except GeneratorExit:
|
||||
self._finish_attempt(
|
||||
attempt_id, "stopped", attempt_started, "cancelled", len(output)
|
||||
)
|
||||
self._finish_request(call, "stopped", "cancelled", len(output))
|
||||
raise
|
||||
except ProviderFailure as exc:
|
||||
self._finish_attempt(
|
||||
attempt_id, "failed", attempt_started, exc.code, len(output)
|
||||
)
|
||||
self._finish_request(call, "failed", exc.code, len(output))
|
||||
raise LLMGatewayError(
|
||||
"stream_interrupted", "智能解读连接中断,请稍后重试。", partial=True
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
self._finish_attempt(
|
||||
attempt_id, "failed", attempt_started, "unavailable", len(output)
|
||||
)
|
||||
self._finish_request(call, "failed", "unavailable", len(output))
|
||||
raise LLMGatewayError(
|
||||
"stream_interrupted", "智能解读连接中断,请稍后重试。", partial=True
|
||||
) from exc
|
||||
self._finish_attempt(attempt_id, "success", attempt_started, "", len(output))
|
||||
self._finish_request(call, "success", "", len(output))
|
||||
yield LLMStreamEvent("done", request_id=call.request_id)
|
||||
return
|
||||
self._finish_request(call, "failed", last_error, 0)
|
||||
code, message = _safe_error(last_error)
|
||||
raise LLMGatewayError(code, message)
|
||||
|
||||
def _finish_attempt(
|
||||
self,
|
||||
attempt_id: int,
|
||||
status: str,
|
||||
started: float,
|
||||
error_type: str,
|
||||
output_chars: int,
|
||||
) -> None:
|
||||
with self._database.transaction() as connection:
|
||||
self._repository.finish_attempt(
|
||||
connection,
|
||||
attempt_id,
|
||||
status,
|
||||
_now(),
|
||||
round((time.perf_counter() - started) * 1000),
|
||||
error_type,
|
||||
output_chars,
|
||||
)
|
||||
|
||||
def _finish_request(
|
||||
self, call: LLMCall, status: str, error_type: str, output_chars: int
|
||||
) -> None:
|
||||
now = datetime.now(SHANGHAI)
|
||||
with self._database.transaction() as connection:
|
||||
current = connection.execute(
|
||||
"SELECT status FROM llm_requests WHERE id = ?", (call.request_id,)
|
||||
).fetchone()
|
||||
if current is None or str(current["status"]) in {"success", "failed", "stopped"}:
|
||||
return
|
||||
self._repository.set_request_status(
|
||||
connection,
|
||||
call.request_id,
|
||||
status,
|
||||
completed_at=now.isoformat(timespec="seconds"),
|
||||
duration_ms=round((now - call.started_at).total_seconds() * 1000),
|
||||
error_type=error_type,
|
||||
output_chars=output_chars,
|
||||
)
|
||||
if status == "success" and not call.quota_exempt:
|
||||
self._repository.increment_usage(
|
||||
connection, call.user_id, call.usage_date, now.isoformat(timespec="seconds")
|
||||
)
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(SHANGHAI).isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def _safe_error(error_type: str) -> tuple[str, str]:
|
||||
messages = {
|
||||
"capacity": ("model_capacity", "智能解读服务当前繁忙,请稍后重试。"),
|
||||
"rate_limited": ("model_rate_limited", "智能解读请求过于频繁,请稍后重试。"),
|
||||
"authentication": ("model_authentication", "智能解读服务配置失效,请联系管理员。"),
|
||||
"timeout": ("model_timeout", "智能解读等待超时,请稍后重试。"),
|
||||
"network": ("model_network", "智能解读网络暂不可用,请稍后重试。"),
|
||||
}
|
||||
return messages.get(error_type, ("model_unavailable", "智能解读服务暂不可用,请稍后重试。"))
|
||||
Reference in New Issue
Block a user