116 lines
3.9 KiB
Python
116 lines
3.9 KiB
Python
import logging
|
|
import os
|
|
import threading
|
|
import time
|
|
from dataclasses import dataclass
|
|
|
|
import httpx
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ChatModelConfig:
|
|
api_base_url: str
|
|
api_key: str
|
|
model: str
|
|
max_tokens: int
|
|
timeout_seconds: float
|
|
vision_model: str
|
|
ocr_model: str
|
|
vision_max_tokens: int
|
|
vision_timeout_seconds: float
|
|
source: str
|
|
|
|
|
|
_cache_lock = threading.Lock()
|
|
_cached_config: ChatModelConfig | None = None
|
|
_cache_expires_at = 0.0
|
|
|
|
|
|
def _environment_config() -> ChatModelConfig:
|
|
return ChatModelConfig(
|
|
api_base_url=os.getenv(
|
|
"CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
).rstrip("/"),
|
|
api_key=os.getenv("CHAT_API_KEY", ""),
|
|
model=os.getenv("CHAT_MODEL", "qwen-plus"),
|
|
max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))),
|
|
timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))),
|
|
vision_model=os.getenv("VISION_MODEL", "qwen3.6-flash"),
|
|
ocr_model=os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr"),
|
|
vision_max_tokens=max(256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))),
|
|
vision_timeout_seconds=max(10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))),
|
|
source="environment",
|
|
)
|
|
|
|
|
|
def _fetch_runtime_config() -> ChatModelConfig | None:
|
|
url = os.getenv("CHAT_MODEL_CONFIG_URL", "").strip()
|
|
token = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "").strip()
|
|
if not url or not token:
|
|
return None
|
|
response = httpx.get(
|
|
url,
|
|
headers={"X-Avatar-Config-Token": token},
|
|
timeout=max(2.0, float(os.getenv("CHAT_MODEL_CONFIG_TIMEOUT_SECONDS", "5"))),
|
|
)
|
|
response.raise_for_status()
|
|
payload = response.json().get("data") or {}
|
|
api_base_url = str(payload.get("api_base_url") or "").rstrip("/")
|
|
api_key = str(payload.get("api_key") or "")
|
|
model = str(payload.get("model") or "")
|
|
if not api_base_url or not api_key or not model:
|
|
raise ValueError("数字分身专用模型配置不完整")
|
|
return ChatModelConfig(
|
|
api_base_url=api_base_url,
|
|
api_key=api_key,
|
|
model=model,
|
|
max_tokens=max(128, int(payload.get("max_tokens") or 1024)),
|
|
timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)),
|
|
vision_model=str(
|
|
payload.get("vision_model")
|
|
or os.getenv("VISION_MODEL", "qwen3.6-flash")
|
|
),
|
|
ocr_model=str(
|
|
payload.get("ocr_model")
|
|
or os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr")
|
|
),
|
|
vision_max_tokens=max(
|
|
256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))
|
|
),
|
|
vision_timeout_seconds=max(
|
|
10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))
|
|
),
|
|
source="admin",
|
|
)
|
|
|
|
|
|
def get_chat_model_config(*, force_refresh: bool = False) -> ChatModelConfig:
|
|
global _cached_config, _cache_expires_at
|
|
|
|
now = time.monotonic()
|
|
if not force_refresh and _cached_config is not None and now < _cache_expires_at:
|
|
return _cached_config
|
|
|
|
with _cache_lock:
|
|
now = time.monotonic()
|
|
if not force_refresh and _cached_config is not None and now < _cache_expires_at:
|
|
return _cached_config
|
|
try:
|
|
config = _fetch_runtime_config() or _environment_config()
|
|
except (httpx.HTTPError, ValueError, TypeError) as exc:
|
|
logger.warning("读取数字分身专用模型配置失败,暂时使用环境变量配置: %s", exc)
|
|
config = _environment_config()
|
|
_cached_config = config
|
|
ttl = max(5, int(os.getenv("CHAT_MODEL_CONFIG_CACHE_SECONDS", "60")))
|
|
_cache_expires_at = now + ttl
|
|
return config
|
|
|
|
|
|
def clear_chat_model_config_cache() -> None:
|
|
global _cached_config, _cache_expires_at
|
|
with _cache_lock:
|
|
_cached_config = None
|
|
_cache_expires_at = 0.0
|