feat(ai): add dedicated digital avatar model config
This commit is contained in:
@@ -23,13 +23,10 @@ from services.token_billing import (
|
||||
reserve_avatar_tokens,
|
||||
settle_reservation,
|
||||
)
|
||||
from services.chat_model_config import ChatModelConfig, get_chat_model_config
|
||||
|
||||
router = APIRouter(tags=["数字分身聊天"])
|
||||
|
||||
CHAT_API_URL = os.getenv("CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1")
|
||||
CHAT_API_KEY = os.getenv("CHAT_API_KEY", "")
|
||||
CHAT_MODEL = os.getenv("CHAT_MODEL", "qwen-plus")
|
||||
CHAT_MAX_OUTPUT_TOKENS = max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024")))
|
||||
MAX_MESSAGE_LENGTH = 4000
|
||||
MAX_HISTORY_MESSAGES = 10
|
||||
QA_LEXICAL_THRESHOLD = 0.72
|
||||
@@ -287,22 +284,25 @@ def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5
|
||||
return results
|
||||
|
||||
|
||||
def _call_qwen(messages: list[dict], temperature: float) -> dict:
|
||||
if not CHAT_API_KEY:
|
||||
def _call_qwen(
|
||||
messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None
|
||||
) -> dict:
|
||||
model_config = model_config or get_chat_model_config()
|
||||
if not model_config.api_key:
|
||||
raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY")
|
||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
||||
url = f"{model_config.api_base_url}/chat/completions"
|
||||
payload = {
|
||||
"model": CHAT_MODEL,
|
||||
"model": model_config.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
||||
"max_tokens": model_config.max_tokens,
|
||||
}
|
||||
try:
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {CHAT_API_KEY}"},
|
||||
headers={"Authorization": f"Bearer {model_config.api_key}"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
timeout=model_config.timeout_seconds,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
@@ -314,21 +314,30 @@ def _call_qwen(messages: list[dict], temperature: float) -> dict:
|
||||
return {"answer": answer.strip(), "usage": data.get("usage") or {}}
|
||||
|
||||
|
||||
def _iter_qwen_stream(messages: list[dict], temperature: float):
|
||||
def _iter_qwen_stream(
|
||||
messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None
|
||||
):
|
||||
"""将 OpenAI 兼容接口的 SSE 分片原样转为文本增量。"""
|
||||
if not CHAT_API_KEY:
|
||||
model_config = model_config or get_chat_model_config()
|
||||
if not model_config.api_key:
|
||||
raise RuntimeError("模型服务未配置")
|
||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
||||
url = f"{model_config.api_base_url}/chat/completions"
|
||||
payload = {
|
||||
"model": CHAT_MODEL,
|
||||
"model": model_config.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
||||
"max_tokens": model_config.max_tokens,
|
||||
"stream": True,
|
||||
"stream_options": {"include_usage": True},
|
||||
}
|
||||
try:
|
||||
with httpx.stream("POST", url, headers={"Authorization": f"Bearer {CHAT_API_KEY}"}, json=payload, timeout=45) as response:
|
||||
with httpx.stream(
|
||||
"POST",
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {model_config.api_key}"},
|
||||
json=payload,
|
||||
timeout=max(45, model_config.timeout_seconds),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
for raw_line in response.iter_lines():
|
||||
line = raw_line.decode() if isinstance(raw_line, bytes) else raw_line
|
||||
@@ -387,16 +396,21 @@ def _resolve_reply(
|
||||
if model_client is not None:
|
||||
answer = model_client(messages=messages, temperature=temperature)
|
||||
else:
|
||||
model_config = get_chat_model_config()
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
usage_source,
|
||||
CHAT_MODEL,
|
||||
model_config.model,
|
||||
messages,
|
||||
CHAT_MAX_OUTPUT_TOKENS,
|
||||
model_config.max_tokens,
|
||||
)
|
||||
try:
|
||||
model_result = _call_qwen(messages=messages, temperature=temperature)
|
||||
model_result = _call_qwen(
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
model_config=model_config,
|
||||
)
|
||||
answer = model_result["answer"]
|
||||
token_usage = settle_reservation(
|
||||
db,
|
||||
@@ -436,15 +450,16 @@ def _stream_reply(
|
||||
config = _config(avatar)
|
||||
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||
messages = _build_prompt(avatar, history, question, references)
|
||||
model_config = get_chat_model_config()
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
usage_source,
|
||||
CHAT_MODEL,
|
||||
model_config.model,
|
||||
messages,
|
||||
CHAT_MAX_OUTPUT_TOKENS,
|
||||
model_config.max_tokens,
|
||||
)
|
||||
chunks = _iter_qwen_stream(messages, temperature)
|
||||
chunks = _iter_qwen_stream(messages, temperature, model_config)
|
||||
if matched:
|
||||
messages, reservation = [], None
|
||||
if public:
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
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
|
||||
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"))),
|
||||
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)),
|
||||
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
|
||||
@@ -0,0 +1,79 @@
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
from services.chat_model_config import (
|
||||
clear_chat_model_config_cache,
|
||||
get_chat_model_config,
|
||||
)
|
||||
|
||||
|
||||
def setup_function():
|
||||
clear_chat_model_config_cache()
|
||||
|
||||
|
||||
def teardown_function():
|
||||
clear_chat_model_config_cache()
|
||||
|
||||
|
||||
def test_admin_runtime_config_takes_priority(monkeypatch):
|
||||
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime")
|
||||
monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret")
|
||||
response = Mock()
|
||||
response.raise_for_status.return_value = None
|
||||
response.json.return_value = {
|
||||
"data": {
|
||||
"api_base_url": "https://model.test/v1/",
|
||||
"api_key": "runtime-key",
|
||||
"model": "avatar-model",
|
||||
"max_tokens": 2048,
|
||||
"timeout_seconds": 42,
|
||||
}
|
||||
}
|
||||
|
||||
with patch("services.chat_model_config.httpx.get", return_value=response) as request:
|
||||
config = get_chat_model_config()
|
||||
|
||||
assert config.source == "admin"
|
||||
assert config.api_base_url == "https://model.test/v1"
|
||||
assert config.model == "avatar-model"
|
||||
assert config.max_tokens == 2048
|
||||
request.assert_called_once_with(
|
||||
"http://config.test/runtime",
|
||||
headers={"X-Avatar-Config-Token": "shared-secret"},
|
||||
timeout=5.0,
|
||||
)
|
||||
|
||||
|
||||
def test_runtime_failure_falls_back_to_environment(monkeypatch):
|
||||
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime")
|
||||
monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret")
|
||||
monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/")
|
||||
monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
|
||||
monkeypatch.setenv("CHAT_MODEL", "fallback-model")
|
||||
monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
|
||||
|
||||
request = httpx.Request("GET", "http://config.test/runtime")
|
||||
with patch(
|
||||
"services.chat_model_config.httpx.get",
|
||||
side_effect=httpx.ConnectError("offline", request=request),
|
||||
):
|
||||
config = get_chat_model_config()
|
||||
|
||||
assert config.source == "environment"
|
||||
assert config.api_base_url == "https://fallback.test/v1"
|
||||
assert config.api_key == "fallback-key"
|
||||
assert config.model == "fallback-model"
|
||||
assert config.max_tokens == 1536
|
||||
|
||||
|
||||
def test_runtime_config_is_cached(monkeypatch):
|
||||
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "")
|
||||
monkeypatch.setenv("CHAT_MODEL", "first-model")
|
||||
first = get_chat_model_config()
|
||||
monkeypatch.setenv("CHAT_MODEL", "second-model")
|
||||
|
||||
second = get_chat_model_config()
|
||||
|
||||
assert first is second
|
||||
assert second.model == "first-model"
|
||||
Reference in New Issue
Block a user