feat(ai): add dedicated digital avatar model config

This commit is contained in:
stefanfeng
2026-08-25 16:50:33 +08:00
parent 3f7ff9329a
commit dc34a03357
13 changed files with 305 additions and 32 deletions
+52 -4
View File
@@ -1,8 +1,11 @@
"""AI模型配置接口"""
from fastapi import APIRouter, Depends, HTTPException
import secrets
from fastapi import APIRouter, Depends, Header, HTTPException
from sqlalchemy import select, update
from app.core.database import get_db
from app.core.config import settings
from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest
from app.models import AIModelConfig
from app.utils.crypto import encrypt, decrypt
@@ -22,10 +25,15 @@ async def list_models(db=Depends(get_db)):
@router.post("")
async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
if req.is_default:
await db.execute(update(AIModelConfig).values(is_default=0))
await db.execute(
update(AIModelConfig)
.where(AIModelConfig.usage_scope == req.usage_scope)
.values(is_default=0)
)
model = AIModelConfig(
model_name=req.model_name,
provider=req.provider,
usage_scope=req.usage_scope,
api_base_url=req.api_base_url,
api_key_enc=encrypt(req.api_key) if req.api_key else None,
model_version=req.model_version,
@@ -47,8 +55,16 @@ async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
if req.is_default:
await db.execute(update(AIModelConfig).where(AIModelConfig.id != model_id).values(is_default=0))
target_scope = req.usage_scope or model.usage_scope
if req.is_default or (req.usage_scope and model.is_default):
await db.execute(
update(AIModelConfig)
.where(
AIModelConfig.id != model_id,
AIModelConfig.usage_scope == target_scope,
)
.values(is_default=0)
)
for field, val in req.model_dump(exclude_none=True).items():
if field == "api_key":
model.api_key_enc = encrypt(val) if val else None
@@ -59,6 +75,37 @@ async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_
return ApiResponse(data=_format_model(model), message="更新成功")
@router.get("/runtime/digital-avatar")
async def get_digital_avatar_runtime_model(
x_avatar_config_token: str | None = Header(default=None),
db=Depends(get_db),
):
expected = settings.AVATAR_MODEL_CONFIG_TOKEN
if not expected:
raise HTTPException(status_code=503, detail="数字分身模型配置服务未启用")
if not x_avatar_config_token or not secrets.compare_digest(x_avatar_config_token, expected):
raise HTTPException(status_code=401, detail="无权读取数字分身模型配置")
result = await db.execute(
select(AIModelConfig).where(
AIModelConfig.usage_scope == "digital_avatar",
AIModelConfig.is_default == 1,
AIModelConfig.is_enabled == 1,
)
)
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="尚未配置启用的数字分身专用模型")
return ApiResponse(data={
"api_base_url": model.api_base_url or "https://api.openai.com/v1",
"api_key": decrypt(model.api_key_enc) if model.api_key_enc else "",
"model": model.model_version or model.model_name,
"temperature": model.temperature,
"max_tokens": model.max_tokens,
"timeout_seconds": model.timeout_seconds,
})
@router.delete("/{model_id}")
async def delete_model(model_id: int, db=Depends(get_db)):
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
@@ -79,6 +126,7 @@ async def test_model(req: AIModelTestRequest, db=Depends(get_db)):
def _format_model(m: AIModelConfig) -> dict:
return {
"id": m.id, "model_name": m.model_name, "provider": m.provider,
"usage_scope": m.usage_scope,
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
"model_version": m.model_version, "temperature": m.temperature,
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
+1
View File
@@ -19,6 +19,7 @@ class Settings(BaseSettings):
# 安全
SECRET_KEY: str = os.getenv("SECRET_KEY", "dev-secret-key-change-in-prod")
AES_KEY: str = os.getenv("AES_KEY", "your-aes-key-32-chars-change-now!")
AVATAR_MODEL_CONFIG_TOKEN: str = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "")
# 新闻平台
NEWS_PLATFORM_BASE_URL: str = os.getenv(
+13
View File
@@ -1,6 +1,7 @@
"""数据库连接管理"""
import asyncio
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy import text
from sqlalchemy.orm import DeclarativeBase
from app.core.config import settings
from app.core.logger import logger
@@ -64,6 +65,18 @@ async def init_db():
VirtualUser, UserPersonality, InteractionRecord,
PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog
)
async with engine.begin() as conn:
result = await conn.execute(text(
"SELECT COUNT(*) FROM information_schema.COLUMNS "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
"AND COLUMN_NAME = 'usage_scope'"
))
if result.scalar_one() == 0:
await conn.execute(text(
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope "
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider"
))
logger.info("AI模型配置表已增加 usage_scope 字段")
logger.info("✅ 数据库模型注册成功")
logger.info("✅ 数据库初始化完成")
+1
View File
@@ -122,6 +122,7 @@ class AIModelConfig(Base):
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
model_name: Mapped[str] = mapped_column(String(64), nullable=False)
provider: Mapped[str] = mapped_column(String(32), nullable=False)
usage_scope: Mapped[str] = mapped_column(String(16), nullable=False, default="general")
api_base_url: Mapped[str | None] = mapped_column(String(256))
api_key_enc: Mapped[str | None] = mapped_column(String(512))
model_version: Mapped[str | None] = mapped_column(String(64))
+4
View File
@@ -154,6 +154,7 @@ class InteractionResponse(BaseModel):
class AIModelCreateRequest(BaseModel):
model_name: str = Field(..., min_length=1, max_length=64)
provider: str = Field(..., pattern="^(openai|zhipu|wenxin|qianwen|local)$")
usage_scope: str = Field(default="general", pattern="^(general|digital_avatar)$")
api_base_url: Optional[str] = None
api_key: Optional[str] = None
model_version: Optional[str] = None
@@ -165,6 +166,8 @@ class AIModelCreateRequest(BaseModel):
class AIModelUpdateRequest(BaseModel):
model_name: Optional[str] = None
provider: Optional[str] = Field(None, pattern="^(openai|zhipu|wenxin|qianwen|local)$")
usage_scope: Optional[str] = Field(None, pattern="^(general|digital_avatar)$")
api_base_url: Optional[str] = None
api_key: Optional[str] = None
model_version: Optional[str] = None
@@ -179,6 +182,7 @@ class AIModelResponse(BaseModel):
id: int
model_name: str
provider: str
usage_scope: str
api_base_url: Optional[str]
has_api_key: bool
model_version: Optional[str]
+3 -1
View File
@@ -28,7 +28,9 @@ class AIService:
async def _get_default_model(self, db: AsyncSession) -> Optional[AIModelConfig]:
result = await db.execute(
select(AIModelConfig).where(
AIModelConfig.is_default == 1, AIModelConfig.is_enabled == 1
AIModelConfig.usage_scope == "general",
AIModelConfig.is_default == 1,
AIModelConfig.is_enabled == 1,
)
)
return result.scalar_one_or_none()
+38 -23
View File
@@ -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"
+3
View File
@@ -10,6 +10,9 @@ services:
environment:
DATABASE_URL: sqlite:////data/avatar.db
UPLOAD_DIR: /data/uploads
CHAT_MODEL_CONFIG_URL: http://host.docker.internal:8000/api/ai-models/runtime/digital-avatar
extra_hosts:
- "host.docker.internal:host-gateway"
volumes:
- avatar-data:/data
expose:
+1
View File
@@ -17,6 +17,7 @@ services:
- REDIS_PORT=6379
- SECRET_KEY=your-secret-key-change-in-production
- AES_KEY=your-aes-key-32-chars-change-now!
- AVATAR_MODEL_CONFIG_TOKEN=${AVATAR_MODEL_CONFIG_TOKEN:-}
- TZ=Asia/Shanghai
- AVATAR_DB_PATH=/app/avatar.db
volumes:
+1
View File
@@ -94,6 +94,7 @@ CREATE TABLE IF NOT EXISTS `ai_model_configs` (
`id` bigint NOT NULL AUTO_INCREMENT,
`model_name` varchar(64) NOT NULL COMMENT '模型名称',
`provider` varchar(32) NOT NULL COMMENT 'openai/zhipu/wenxin/qianwen/local',
`usage_scope` varchar(16) NOT NULL DEFAULT 'general' COMMENT '用途:general/digital_avatar',
`api_base_url` varchar(256) DEFAULT NULL COMMENT 'API地址',
`api_key_enc` varchar(512) DEFAULT NULL COMMENT '加密API Key',
`model_version` varchar(64) DEFAULT NULL COMMENT '模型版本',
+16 -4
View File
@@ -15,6 +15,9 @@
<span class="model-title">{{ m.model_name }}</span>
</div>
<div style="display:flex;gap:6px;align-items:center">
<el-tag :type="m.usage_scope === 'digital_avatar' ? 'warning' : 'info'" size="small">
{{ scopeLabels[m.usage_scope] || '通用业务' }}
</el-tag>
<el-tag v-if="m.is_default" type="success" size="small">默认</el-tag>
<el-tag v-if="!m.is_enabled" type="danger" size="small">禁用</el-tag>
</div>
@@ -49,6 +52,13 @@
<el-option v-for="(l,v) in providerLabels" :key="v" :label="l" :value="v" />
</el-select>
</el-form-item>
<el-form-item label="使用场景" prop="usage_scope">
<el-radio-group v-model="form.usage_scope">
<el-radio-button value="general">通用业务</el-radio-button>
<el-radio-button value="digital_avatar">数字分身专用</el-radio-button>
</el-radio-group>
<div class="scope-tip">数字分身专用模型仅用于分身对话和主动接管回复</div>
</el-form-item>
<el-form-item label="API地址">
<el-input v-model="form.api_base_url" placeholder="留空使用默认地址" />
</el-form-item>
@@ -131,8 +141,9 @@ const testResult = ref(null)
const testing = ref(false)
const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' }
const form = reactive({ model_name: '', provider: 'openai', api_base_url: '', api_key: '', model_version: '', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }] }
const scopeLabels = { general: '通用业务', digital_avatar: '数字分身专用' }
const form = reactive({ model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: '', api_key: '', model_version: '', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }], usage_scope: [{ required: true }] }
async function load() {
const res = await getAIModels()
@@ -155,13 +166,13 @@ function onProviderChange(provider) {
function openCreate() {
editModel.value = null
Object.assign(form, { model_name: '', provider: 'openai', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
Object.assign(form, { model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
dialogVisible.value = true
}
function openEdit(m) {
editModel.value = m
Object.assign(form, { model_name: m.model_name, provider: m.provider, api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default })
Object.assign(form, { model_name: m.model_name, provider: m.provider, usage_scope: m.usage_scope || 'general', api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default })
dialogVisible.value = true
}
@@ -236,4 +247,5 @@ onMounted(load)
.result-meta { display: flex; align-items: center; gap: 10px; margin-bottom: 10px; }
.result-content { background: var(--color-bg); border: 1px solid var(--color-border); border-radius: 8px; padding: 12px; font-size: 13px; line-height: 1.6; white-space: pre-wrap; max-height: 200px; overflow-y: auto; }
.empty-state { grid-column: 1/-1; padding: 40px; }
.scope-tip { margin-top: 6px; color: var(--color-text-muted); font-size: 12px; line-height: 1.5; }
</style>