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
+2 -1
View File
@@ -1,7 +1,8 @@
from .m0001_adopt_legacy import MIGRATION as M0001_ADOPT_LEGACY from .m0001_adopt_legacy import MIGRATION as M0001_ADOPT_LEGACY
from .m0002_job_runs import MIGRATION as M0002_JOB_RUNS from .m0002_job_runs import MIGRATION as M0002_JOB_RUNS
from .m0003_llm_audit import MIGRATION as M0003_LLM_AUDIT
from .runner import Migration, MigrationError, MigrationRunner from .runner import Migration, MigrationError, MigrationRunner
MIGRATIONS = (M0001_ADOPT_LEGACY, M0002_JOB_RUNS) MIGRATIONS = (M0001_ADOPT_LEGACY, M0002_JOB_RUNS, M0003_LLM_AUDIT)
__all__ = ["MIGRATIONS", "Migration", "MigrationError", "MigrationRunner"] __all__ = ["MIGRATIONS", "Migration", "MigrationError", "MigrationRunner"]
@@ -0,0 +1,32 @@
from __future__ import annotations
import sqlite3
from backend.database.migrations.runner import Migration
def extend_llm_audit(connection: sqlite3.Connection) -> None:
columns = {
str(row["name"])
for row in connection.execute("PRAGMA table_info(llm_usage)")
}
additions = (
("role", "TEXT NOT NULL DEFAULT ''"),
("prompt_version", "TEXT NOT NULL DEFAULT ''"),
("error_code", "TEXT NOT NULL DEFAULT ''"),
("input_tokens", "INTEGER NOT NULL DEFAULT 0"),
("output_tokens", "INTEGER NOT NULL DEFAULT 0"),
)
for name, declaration in additions:
if name not in columns:
connection.execute(
f"ALTER TABLE llm_usage ADD COLUMN {name} {declaration}"
)
MIGRATION = Migration(
version="0003",
name="extend_llm_audit",
action=extend_llm_audit,
signature="llm-audit:v1:role,prompt-version,error-code,input-tokens,output-tokens",
)
+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]
+13 -3
View File
@@ -874,16 +874,26 @@ class ReviewDatabase:
model: str, model: str,
status: str, status: str,
latency_ms: int = 0, latency_ms: int = 0,
*,
role: str = "",
prompt_version: str = "",
error_code: str = "",
input_tokens: int = 0,
output_tokens: int = 0,
) -> None: ) -> None:
now = datetime.now(timezone.utc).isoformat(timespec="seconds") now = datetime.now(timezone.utc).isoformat(timespec="seconds")
with self.connect() as connection: with self.connect() as connection:
connection.execute( connection.execute(
""" """
INSERT INTO llm_usage INSERT INTO llm_usage
(user_id, feature, source, model, status, latency_ms, created_at) (user_id, feature, source, model, status, latency_ms, created_at,
VALUES (?, ?, ?, ?, ?, ?, ?) role, prompt_version, error_code, input_tokens, output_tokens)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""", """,
(user_id, feature, source, model, status, int(latency_ms), now), (
user_id, feature, source, model, status, int(latency_ms), now,
role, prompt_version, error_code, int(input_tokens), int(output_tokens),
),
) )
def count_llm_usage_since(self, user_id: int, source: str, since: str) -> int: def count_llm_usage_since(self, user_id: int, source: str, since: str) -> int:
+7 -7
View File
@@ -262,16 +262,16 @@
"bytes": 364385, "bytes": 364385,
"lines": 15549 "lines": 15549
}, },
{
"path": "server.py",
"bytes": 268447,
"lines": 5956
},
{ {
"path": "static/redesign-v2.css", "path": "static/redesign-v2.css",
"bytes": 264960, "bytes": 264960,
"lines": 8616 "lines": 8616
}, },
{
"path": "server.py",
"bytes": 263744,
"lines": 5856
},
{ {
"path": "static/index.html", "path": "static/index.html",
"bytes": 133569, "bytes": 133569,
@@ -279,8 +279,8 @@
}, },
{ {
"path": "database.py", "path": "database.py",
"bytes": 121153, "bytes": 121546,
"lines": 2829 "lines": 2839
}, },
{ {
"path": "screener.py", "path": "screener.py",
+32
View File
@@ -0,0 +1,32 @@
# Stage 13: Unified LLM Gateway
## Boundary
All runtime model calls now enter through `backend/llm/gateway.py`. Feature agents retain
their deterministic context assembly, prompt content, and provider response parsing.
The gateway owns:
- membership and daily quota enforcement;
- primary and fallback model selection;
- fallback only before the first streamed delta;
- stable user-visible availability and interruption errors;
- one logical-call audit record with feature, model role, prompt version, latency, status,
normalized error code, and token fields reserved for providers that report usage.
Administrator connection probes also cross the gateway boundary, but do not consume member
quota or create usage records.
## Compatibility
- Mentor and review-assistant stream payloads are unchanged.
- Heaven readings retain their primary/fallback notice and persistence behavior.
- Strategy compilation still falls back to the deterministic local compiler when both model
profiles are unavailable, while access and quota failures remain blocking.
- Provider credentials and raw provider failures remain outside browser responses.
## Residual Risk
The current OpenAI-compatible streaming providers do not consistently return token usage, so
the audit schema records zero until transport adapters expose trustworthy token counts. Request
cancellation remains bounded by the existing provider socket timeout.
+112 -212
View File
@@ -19,6 +19,7 @@ from assistant_agent import ReviewAssistantError, stream_review_assistant
from api_access import ROUTES from api_access import ROUTES
from backend.bootstrap import build_application_container, load_runtime_settings from backend.bootstrap import build_application_container, load_runtime_settings
from backend.http import correlation_id, normalize_error_payload from backend.http import correlation_id, normalize_error_payload
from backend.llm import LLMGateway, LLMGatewayError
from chart_data_provider import ChartDataError from chart_data_provider import ChartDataError
from app_config import ( from app_config import (
DATA_DIR, DATA_DIR,
@@ -193,6 +194,13 @@ class DashboardService:
self.realtime_aggregator = self.container.realtime_aggregator self.realtime_aggregator = self.container.realtime_aggregator
self.chart_data = self.container.chart_data self.chart_data = self.container.chart_data
self.jobs = self.container.jobs self.jobs = self.container.jobs
self.llm_gateway = LLMGateway(
database=self.database,
user_id_supplier=lambda: self.current_user_id,
membership_supplier=self.membership,
settings_supplier=lambda: self._system_credentials,
profile_supplier=self._resolved_llm_profile,
)
self.screener.ensure_builtin_strategies() self.screener.ensure_builtin_strategies()
self._background_stop = threading.Event() self._background_stop = threading.Event()
self._background_thread = self.jobs.start_scheduler( self._background_thread = self.jobs.start_scheduler(
@@ -482,7 +490,12 @@ class DashboardService:
raise ValueError("模型角色不支持。") raise ValueError("模型角色不支持。")
profile = self._validate_llm_profile(payload, current, required=True, label=label) profile = self._validate_llm_profile(payload, current, required=True, label=label)
try: try:
return test_llm_connection(**profile) return self.llm_gateway.probe(
profile,
lambda model: test_llm_connection(
model.api_key, model.base_url, model.model
),
)
except LLMCompilerError as exc: except LLMCompilerError as exc:
raise ValueError(str(exc)) from exc raise ValueError(str(exc)) from exc
@@ -533,29 +546,6 @@ class DashboardService:
start.isoformat(timespec="seconds"), start.isoformat(timespec="seconds"),
) )
def ensure_llm_access(self, feature: str) -> str:
profile = self._resolved_llm_profile()
source = str(profile.get("source") or "none")
if source == "none" or not self._profile_configured(profile.get("primary") or {}):
raise ValueError("智能功能尚未配置,请联系管理员。")
if source == "platform":
limit = max(1, int(self._system_credentials.get("member_daily_limit") or 50))
if self._platform_usage_today() >= limit:
raise ValueError(f"今日会员模型额度已用完({limit} 次)。")
return source
def record_llm_usage(
self,
feature: str,
source: str,
model: str,
status: str,
latency_ms: int = 0,
) -> None:
self.database.record_llm_usage(
self.current_user_id, feature, source, model, status, latency_ms
)
def system_status(self) -> dict[str, Any]: def system_status(self) -> dict[str, Any]:
platform = self._platform_llm_profile() platform = self._platform_llm_profile()
model_pool = [] model_pool = []
@@ -707,7 +697,12 @@ class DashboardService:
payload, current, required=True, label=label payload, current, required=True, label=label
) )
try: try:
return test_llm_connection(**profile) return self.llm_gateway.probe(
profile,
lambda model: test_llm_connection(
model.api_key, model.base_url, model.model
),
)
except LLMCompilerError as exc: except LLMCompilerError as exc:
raise ValueError(str(exc)) from exc raise ValueError(str(exc)) from exc
@@ -1617,55 +1612,33 @@ class DashboardService:
for item in self.assistant_messages()[-12:] for item in self.assistant_messages()[-12:]
if item.get("role") in {"user", "assistant"} if item.get("role") in {"user", "assistant"}
] ]
source = self.ensure_llm_access("assistant")
profiles = [
(
self.llm_primary_api_key,
self.llm_primary_base_url,
self.llm_primary_model,
)
]
if self.llm_fallback_configured:
profiles.append(
(
self.llm_fallback_api_key,
self.llm_fallback_base_url,
self.llm_fallback_model,
)
)
def generate(): def generate():
started = time.perf_counter() answer_parts: list[str] = []
last_error: Exception | None = None events = self.llm_gateway.stream(
for api_key, base_url, model in profiles: "assistant",
try: "review-assistant-v1",
upstream = iter( lambda profile: stream_review_assistant(
stream_review_assistant( context,
context, question, history, api_key, base_url, model question,
) history,
profile.api_key,
profile.base_url,
profile.model,
),
(ReviewAssistantError,),
)
for event in events:
if event.kind == "delta":
chunk = str(event.value or "")
answer_parts.append(chunk)
yield chunk
elif event.kind == "complete":
self.database.save_assistant_exchange(
self.current_user_id,
question,
"".join(answer_parts).strip(),
trade_date,
) )
first = next(upstream)
except (ReviewAssistantError, StopIteration) as exc:
last_error = exc
continue
answer_parts = [first]
yield first
try:
for chunk in upstream:
answer_parts.append(chunk)
yield chunk
except ReviewAssistantError as exc:
self.record_llm_usage("assistant", source, model, "failed")
raise ValueError("智能解读连接中断,请稍后重试。") from exc
answer = "".join(answer_parts).strip()
latency_ms = round((time.perf_counter() - started) * 1000)
self.database.save_assistant_exchange(
self.current_user_id, question, answer, trade_date
)
self.record_llm_usage("assistant", source, model, "success", latency_ms)
return
self.record_llm_usage("assistant", source, self.llm_primary_model, "failed")
raise ValueError("智能解读服务暂不可用,请稍后重试。") from last_error
return generate() return generate()
@@ -1867,53 +1840,35 @@ class DashboardService:
if regime not in REGIMES: if regime not in REGIMES:
raise ValueError("市场阶段不支持。") raise ValueError("市场阶段不支持。")
notice = "" notice = ""
compiled = None
primary_error = ""
source = self.llm_source source = self.llm_source
started = datetime.now(timezone.utc)
if source == "platform": if source == "platform":
self.ensure_llm_access("screener")
if self.llm_configured:
try: try:
compiled = compile_strategy_with_llm( gateway_result = self.llm_gateway.call(
prompt, "screener",
regime, "strategy-compiler-v1",
self.llm_primary_api_key, lambda profile: compile_strategy_with_llm(
self.llm_primary_base_url, prompt,
self.llm_primary_model, regime,
profile.api_key,
profile.base_url,
profile.model,
),
(LLMCompilerError,),
) )
except LLMCompilerError as exc: compiled = gateway_result.value
primary_error = str(exc) if gateway_result.role == "fallback":
if compiled is None and self.llm_fallback_configured: compiled["compiler"] = "llm_fallback"
try: notice = "智能策略生成服务已自动切换。"
compiled = compile_strategy_with_llm( except LLMGatewayError as exc:
prompt, if exc.code != "unavailable":
regime, raise
self.llm_fallback_api_key,
self.llm_fallback_base_url,
self.llm_fallback_model,
)
compiled["compiler"] = "llm_fallback"
notice = "智能策略生成服务已自动切换。"
except LLMCompilerError as exc:
fallback_error = str(exc)
compiled = compile_local_strategy(prompt, regime) compiled = compile_local_strategy(prompt, regime)
notice = "智能策略生成暂不可用,已使用本地模板。" notice = "智能策略生成暂不可用,已使用本地模板。"
if compiled is None: else:
compiled = compile_local_strategy(prompt, regime) compiled = compile_local_strategy(prompt, regime)
notice = "智能策略生成暂不可用,已使用本地模板。" notice = "智能策略生成暂不可用,已使用本地模板。"
compiled["formula"] = self.screener.validate_formula(compiled["formula"]) compiled["formula"] = self.screener.validate_formula(compiled["formula"])
compiled["notice"] = notice compiled["notice"] = notice
if source in {"personal", "platform"}:
elapsed = int((datetime.now(timezone.utc) - started).total_seconds() * 1000)
status = "success" if str(compiled.get("compiler") or "").startswith("llm") else "failed"
self.record_llm_usage(
"screener",
source,
str(compiled.get("model") or self.llm_primary_model),
status,
elapsed,
)
return compiled return compiled
def save_screener_strategy(self, payload: dict[str, Any]) -> dict[str, Any]: def save_screener_strategy(self, payload: dict[str, Any]) -> dict[str, Any]:
@@ -2017,72 +1972,43 @@ class DashboardService:
) )
context = self._build_mentor_context(trade_date, question, skill) context = self._build_mentor_context(trade_date, question, skill)
source = self.ensure_llm_access("mentor")
profiles = []
if self.llm_configured:
profiles.append(
(
"primary",
self.llm_primary_api_key,
self.llm_primary_base_url,
self.llm_primary_model,
)
)
if self.llm_fallback_configured:
profiles.append(
(
"fallback",
self.llm_fallback_api_key,
self.llm_fallback_base_url,
self.llm_fallback_model,
)
)
def generate(): def generate():
started = time.perf_counter() answer_parts: list[str] = []
last_error: Exception | None = None events = self.llm_gateway.stream(
for compiler, api_key, base_url, model in profiles: "mentor",
try: f"mentor-skill-v1:{skill.skill_id}",
upstream = iter( lambda profile: stream_with_mentor(
stream_with_mentor( skill,
skill, context, question, history, api_key, base_url, model context,
)
)
first = next(upstream)
except (MentorAgentError, StopIteration) as exc:
last_error = exc
continue
answer_parts = [first]
yield {"type": "delta", "content": first}
try:
for chunk in upstream:
answer_parts.append(chunk)
yield {"type": "delta", "content": chunk}
except MentorAgentError as exc:
self.record_llm_usage("mentor", source, model, "failed")
raise ValueError("智能解读连接中断,请稍后重试。") from exc
answer = "".join(answer_parts).strip()
latency_ms = round((time.perf_counter() - started) * 1000)
self.database.save_mentor_exchange(
self.current_user_id,
mentor_id,
trade_date,
question, question,
answer, history,
context["data_trade_date"], profile.api_key,
) profile.base_url,
self.record_llm_usage("mentor", source, model, "success", latency_ms) profile.model,
yield { ),
"type": "meta", (MentorAgentError,),
"data_trade_date": context["data_trade_date"], )
"notice": "智能解读已自动切换可用服务。" for event in events:
if compiler == "fallback" if event.kind == "delta":
else "", chunk = str(event.value or "")
} answer_parts.append(chunk)
return yield {"type": "delta", "content": chunk}
failed_model = profiles[-1][3] if profiles else self.llm_primary_model elif event.kind == "complete":
self.record_llm_usage("mentor", source, failed_model, "failed") self.database.save_mentor_exchange(
raise ValueError("智能解读服务暂不可用,请稍后重试。") from last_error self.current_user_id,
mentor_id,
trade_date,
question,
"".join(answer_parts).strip(),
context["data_trade_date"],
)
yield {
"type": "meta",
"data_trade_date": context["data_trade_date"],
"notice": "智能解读已自动切换可用服务。"
if event.role == "fallback"
else "",
}
return generate() return generate()
@@ -3155,45 +3081,19 @@ class DashboardService:
return bool(reading and str(reading.get("answer") or "").rstrip().endswith("……")) return bool(reading and str(reading.get("answer") or "").rstrip().endswith("……"))
def _call_heaven_agent(self, mode: str, context: dict[str, Any]) -> tuple[dict[str, Any], str]: def _call_heaven_agent(self, mode: str, context: dict[str, Any]) -> tuple[dict[str, Any], str]:
source = self.ensure_llm_access(f"heaven_{mode}") result = self.llm_gateway.call(
primary_error = "" f"heaven_{mode}",
if self.llm_configured: f"heaven-{mode}-v1",
try: lambda profile: interpret_heaven(
result = interpret_heaven( mode,
mode, context,
context, profile.api_key,
self.llm_primary_api_key, profile.base_url,
self.llm_primary_base_url, profile.model,
self.llm_primary_model, ),
) (HeavenAgentError,),
self.record_llm_usage( )
f"heaven_{mode}", source, str(result.get("model") or ""), return result.value, result.role
"success", int(result.get("latency_ms") or 0),
)
return result, "primary"
except HeavenAgentError as exc:
primary_error = str(exc)
if self.llm_fallback_configured:
try:
result = interpret_heaven(
mode,
context,
self.llm_fallback_api_key,
self.llm_fallback_base_url,
self.llm_fallback_model,
)
self.record_llm_usage(
f"heaven_{mode}", source, str(result.get("model") or ""),
"success", int(result.get("latency_ms") or 0),
)
return result, "fallback"
except HeavenAgentError as exc:
self.record_llm_usage(
f"heaven_{mode}", source, self.llm_fallback_model, "failed"
)
raise ValueError("智能解读服务暂不可用,请稍后重试。") from exc
self.record_llm_usage(f"heaven_{mode}", source, self.llm_primary_model, "failed")
raise ValueError("智能解读服务暂不可用,请稍后重试。")
def _heaven_index_context( def _heaven_index_context(
self, self,
+2 -1
View File
@@ -23,6 +23,7 @@ class DatabaseMigrationTests(unittest.TestCase):
[ [
("0001", "adopt_legacy_schema"), ("0001", "adopt_legacy_schema"),
("0002", "create_job_runs"), ("0002", "create_job_runs"),
("0003", "extend_llm_audit"),
], ],
) )
ReviewDatabase(path) ReviewDatabase(path)
@@ -30,7 +31,7 @@ class DatabaseMigrationTests(unittest.TestCase):
count = connection.execute( count = connection.execute(
"SELECT COUNT(*) AS count FROM schema_migrations" "SELECT COUNT(*) AS count FROM schema_migrations"
).fetchone()["count"] ).fetchone()["count"]
self.assertEqual(count, 2) self.assertEqual(count, 3)
def test_connection_factory_enables_required_pragmas(self) -> None: def test_connection_factory_enables_required_pragmas(self) -> None:
with tempfile.TemporaryDirectory() as root: with tempfile.TemporaryDirectory() as root:
+118
View File
@@ -0,0 +1,118 @@
from __future__ import annotations
import unittest
from backend.llm import LLMGateway, LLMGatewayError
class ProviderFailure(RuntimeError):
pass
class FakeDatabase:
def __init__(self, used: int = 0) -> None:
self.used = used
self.audit: list[dict[str, object]] = []
def count_llm_usage_since(self, user_id: int, source: str, since: str) -> int:
return self.used
def record_llm_usage(
self,
user_id: int,
feature: str,
source: str,
model: str,
status: str,
latency_ms: int,
**metadata,
) -> None:
self.audit.append(
{
"user_id": user_id,
"feature": feature,
"source": source,
"model": model,
"status": status,
**metadata,
}
)
def profile() -> dict[str, object]:
return {
"source": "platform",
"primary": {"api_key": "p", "base_url": "https://p", "model": "primary"},
"fallback": {"api_key": "f", "base_url": "https://f", "model": "fallback"},
}
class LLMGatewayTests(unittest.TestCase):
def gateway(self, database: FakeDatabase | None = None) -> LLMGateway:
database = database or FakeDatabase()
return LLMGateway(
database=database,
user_id_supplier=lambda: 7,
membership_supplier=lambda: {"active": True},
settings_supplier=lambda: {"member_daily_limit": 50},
profile_supplier=profile,
)
def test_non_streaming_call_falls_back_and_audits_once(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
return {"answer": "ok"}
result = gateway.call("heaven_trend", "heaven-trend-v1", invoke, (ProviderFailure,))
self.assertEqual(result.role, "fallback")
self.assertEqual(result.value, {"answer": "ok"})
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "fallback")
self.assertEqual(database.audit[0]["prompt_version"], "heaven-trend-v1")
def test_stream_falls_back_before_first_delta(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
yield "a"
yield "b"
events = list(gateway.stream("mentor", "mentor-v1", invoke, (ProviderFailure,)))
self.assertEqual([event.value for event in events[:-1]], ["a", "b"])
self.assertEqual(events[-1].kind, "complete")
self.assertEqual(events[-1].role, "fallback")
self.assertEqual(database.audit[0]["status"], "success")
def test_stream_does_not_switch_model_after_output_started(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
yield "first"
raise ProviderFailure(f"{model.role} interrupted")
iterator = gateway.stream("assistant", "assistant-v1", invoke, (ProviderFailure,))
self.assertEqual(next(iterator).value, "first")
with self.assertRaisesRegex(LLMGatewayError, "连接中断"):
list(iterator)
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "primary")
self.assertEqual(database.audit[0]["status"], "failed")
def test_daily_quota_is_enforced_before_provider_call(self) -> None:
gateway = self.gateway(FakeDatabase(used=50))
with self.assertRaisesRegex(LLMGatewayError, "额度已用完"):
gateway.call("mentor", "mentor-v1", lambda model: "unused", (ProviderFailure,))
if __name__ == "__main__":
unittest.main()