233 lines
8.9 KiB
Python
233 lines
8.9 KiB
Python
from __future__ import annotations
|
|
|
|
from datetime import datetime, timezone
|
|
from typing import Any
|
|
from urllib.parse import urlparse
|
|
|
|
from backend.bootstrap.config import validate_text
|
|
from backend.features.screener.compiler import LLMCompilerError, test_llm_connection
|
|
|
|
|
|
class LLMServiceMixin:
|
|
def _personal_llm_profile(self) -> dict[str, Any]:
|
|
credentials = self._credentials()
|
|
return {
|
|
"source": "personal",
|
|
"primary": {
|
|
"api_key": credentials["llm_primary_api_key"],
|
|
"base_url": credentials["llm_primary_base_url"],
|
|
"model": credentials["llm_primary_model"],
|
|
},
|
|
"fallback": {
|
|
"api_key": credentials["llm_fallback_api_key"],
|
|
"base_url": credentials["llm_fallback_base_url"],
|
|
"model": credentials["llm_fallback_model"],
|
|
},
|
|
}
|
|
|
|
def _platform_llm_profile(self) -> dict[str, Any]:
|
|
models = {
|
|
str(item.get("id") or ""): item
|
|
for item in self._system_credentials.get("llm_models") or []
|
|
if isinstance(item, dict) and item.get("id")
|
|
}
|
|
|
|
def selected(role: str) -> dict[str, str]:
|
|
item = models.get(str(self._system_credentials.get(f"{role}_model_id") or ""), {})
|
|
return {
|
|
"id": str(item.get("id") or ""),
|
|
"name": str(item.get("name") or ""),
|
|
"api_key": str(item.get("api_key") or ""),
|
|
"base_url": str(item.get("base_url") or ""),
|
|
"model": str(item.get("model") or ""),
|
|
}
|
|
|
|
return {
|
|
"source": "platform",
|
|
"primary": selected("primary"),
|
|
"fallback": selected("fallback"),
|
|
}
|
|
|
|
@staticmethod
|
|
def _profile_configured(profile: dict[str, str]) -> bool:
|
|
return bool(profile.get("api_key") and profile.get("base_url") and profile.get("model"))
|
|
|
|
def _resolved_llm_profile(self) -> dict[str, Any]:
|
|
platform = self._platform_llm_profile()
|
|
platform_ready = self.membership()["active"] and self._profile_configured(platform["primary"])
|
|
if platform_ready:
|
|
return platform
|
|
return {"source": "none", "primary": {}, "fallback": {}}
|
|
|
|
@property
|
|
def llm_primary_api_key(self) -> str:
|
|
return str(self._resolved_llm_profile()["primary"].get("api_key") or "")
|
|
|
|
@property
|
|
def llm_primary_base_url(self) -> str:
|
|
return str(self._resolved_llm_profile()["primary"].get("base_url") or "")
|
|
|
|
@property
|
|
def llm_primary_model(self) -> str:
|
|
return str(self._resolved_llm_profile()["primary"].get("model") or "")
|
|
|
|
@property
|
|
def llm_fallback_api_key(self) -> str:
|
|
return str(self._resolved_llm_profile()["fallback"].get("api_key") or "")
|
|
|
|
@property
|
|
def llm_fallback_base_url(self) -> str:
|
|
return str(self._resolved_llm_profile()["fallback"].get("base_url") or "")
|
|
|
|
@property
|
|
def llm_fallback_model(self) -> str:
|
|
return str(self._resolved_llm_profile()["fallback"].get("model") or "")
|
|
|
|
@property
|
|
def llm_source(self) -> str:
|
|
return str(self._resolved_llm_profile().get("source") or "none")
|
|
|
|
@property
|
|
def llm_configured(self) -> bool:
|
|
return bool(self.llm_primary_api_key and self.llm_primary_model)
|
|
|
|
@property
|
|
def llm_fallback_configured(self) -> bool:
|
|
return bool(
|
|
self.llm_fallback_api_key
|
|
and self.llm_fallback_base_url
|
|
and self.llm_fallback_model
|
|
)
|
|
|
|
def save_llm_settings(
|
|
self,
|
|
primary: dict[str, Any],
|
|
fallback: dict[str, Any],
|
|
fallback_enabled: bool,
|
|
) -> None:
|
|
personal = self._personal_llm_profile()
|
|
primary_profile = self._validate_llm_profile(
|
|
primary,
|
|
personal["primary"],
|
|
required=True,
|
|
label="主模型",
|
|
)
|
|
if fallback_enabled:
|
|
fallback_profile = self._validate_llm_profile(
|
|
fallback,
|
|
personal["fallback"],
|
|
required=True,
|
|
label="辅助模型",
|
|
)
|
|
else:
|
|
fallback_profile = {"api_key": "", "base_url": "", "model": ""}
|
|
credentials = self._credentials()
|
|
credentials.update(
|
|
{
|
|
"llm_primary_api_key": primary_profile["api_key"],
|
|
"llm_primary_base_url": primary_profile["base_url"],
|
|
"llm_primary_model": primary_profile["model"],
|
|
"llm_fallback_api_key": fallback_profile["api_key"],
|
|
"llm_fallback_base_url": fallback_profile["base_url"],
|
|
"llm_fallback_model": fallback_profile["model"],
|
|
}
|
|
)
|
|
self._save_credentials(credentials)
|
|
|
|
def save_llm_mode(self, mode: str) -> None:
|
|
raise ValueError("LLM 算力由管理员统一配置,会员账号自动使用平台模型。")
|
|
|
|
def test_llm_profile(self, role: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
personal = self._personal_llm_profile()
|
|
if role == "primary":
|
|
current = personal["primary"]
|
|
label = "主模型"
|
|
elif role == "fallback":
|
|
current = personal["fallback"]
|
|
label = "辅助模型"
|
|
else:
|
|
raise ValueError("模型角色不支持。")
|
|
profile = self._validate_llm_profile(payload, current, required=True, label=label)
|
|
try:
|
|
return self.llm_gateway.probe(
|
|
profile,
|
|
lambda model: test_llm_connection(
|
|
model.api_key, model.base_url, model.model
|
|
),
|
|
)
|
|
except LLMCompilerError as exc:
|
|
raise ValueError(str(exc)) from exc
|
|
|
|
@staticmethod
|
|
def _validate_llm_profile(
|
|
payload: dict[str, Any],
|
|
current: dict[str, str],
|
|
required: bool,
|
|
label: str,
|
|
) -> dict[str, str]:
|
|
api_key = str(payload.get("api_key") or current.get("api_key") or "").strip()
|
|
base_url = str(payload.get("base_url") or current.get("base_url") or "").strip().rstrip("/")
|
|
model = str(payload.get("model") or current.get("model") or "").strip()
|
|
if not required and not any((api_key, base_url, model)):
|
|
return {"api_key": "", "base_url": "", "model": ""}
|
|
parsed = urlparse(base_url)
|
|
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
|
raise ValueError(f"{label} Base URL 格式不正确。")
|
|
if not api_key or len(api_key) > 300:
|
|
raise ValueError(f"{label} API Key 不能为空或过长。")
|
|
if not model or len(model) > 100:
|
|
raise ValueError(f"{label}模型名称不能为空或过长。")
|
|
return {"api_key": api_key, "base_url": base_url, "model": model}
|
|
|
|
def llm_access_status(self) -> dict[str, Any]:
|
|
platform = self._platform_llm_profile()
|
|
membership = self.membership()
|
|
limit = max(1, int(self._system_credentials.get("member_daily_limit") or 50))
|
|
used = self._platform_usage_today() if membership["active"] else 0
|
|
resolved = self._resolved_llm_profile()
|
|
return {
|
|
"mode": "platform" if membership["active"] else "locked",
|
|
"resolved_source": resolved.get("source") or "none",
|
|
"resolved_model": str(resolved.get("primary", {}).get("model") or ""),
|
|
"platform_configured": self._profile_configured(platform["primary"]),
|
|
"membership": membership,
|
|
"daily_limit": limit,
|
|
"used_today": used,
|
|
"remaining_calls": None if membership["is_admin"] else max(0, limit - used),
|
|
}
|
|
|
|
def _platform_usage_today(self) -> int:
|
|
return self._platform_usage_today_for_user(self.current_user_id)
|
|
|
|
def _platform_usage_today_for_user(self, user_id: int) -> 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(
|
|
user_id,
|
|
"platform",
|
|
start.isoformat(timespec="seconds"),
|
|
)
|
|
|
|
def test_system_llm_profile(self, model_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
current = next(
|
|
(
|
|
item
|
|
for item in self._system_credentials.get("llm_models") or []
|
|
if str(item.get("id") or "") == model_id
|
|
),
|
|
{},
|
|
)
|
|
label = validate_text(payload.get("name") or current.get("name"), "模型名称", 50, required=True)
|
|
profile = self._validate_llm_profile(
|
|
payload, current, required=True, label=label
|
|
)
|
|
try:
|
|
return self.llm_gateway.probe(
|
|
profile,
|
|
lambda model: test_llm_connection(
|
|
model.api_key, model.base_url, model.model
|
|
),
|
|
)
|
|
except LLMCompilerError as exc:
|
|
raise ValueError(str(exc)) from exc
|