feat(ai): add dedicated digital avatar model config
This commit is contained in:
@@ -1,8 +1,11 @@
|
|||||||
"""AI模型配置接口"""
|
"""AI模型配置接口"""
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
import secrets
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||||
from sqlalchemy import select, update
|
from sqlalchemy import select, update
|
||||||
|
|
||||||
from app.core.database import get_db
|
from app.core.database import get_db
|
||||||
|
from app.core.config import settings
|
||||||
from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest
|
from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest
|
||||||
from app.models import AIModelConfig
|
from app.models import AIModelConfig
|
||||||
from app.utils.crypto import encrypt, decrypt
|
from app.utils.crypto import encrypt, decrypt
|
||||||
@@ -22,10 +25,15 @@ async def list_models(db=Depends(get_db)):
|
|||||||
@router.post("")
|
@router.post("")
|
||||||
async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
|
async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
|
||||||
if req.is_default:
|
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 = AIModelConfig(
|
||||||
model_name=req.model_name,
|
model_name=req.model_name,
|
||||||
provider=req.provider,
|
provider=req.provider,
|
||||||
|
usage_scope=req.usage_scope,
|
||||||
api_base_url=req.api_base_url,
|
api_base_url=req.api_base_url,
|
||||||
api_key_enc=encrypt(req.api_key) if req.api_key else None,
|
api_key_enc=encrypt(req.api_key) if req.api_key else None,
|
||||||
model_version=req.model_version,
|
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()
|
model = result.scalar_one_or_none()
|
||||||
if not model:
|
if not model:
|
||||||
raise HTTPException(status_code=404, detail="模型不存在")
|
raise HTTPException(status_code=404, detail="模型不存在")
|
||||||
if req.is_default:
|
target_scope = req.usage_scope or model.usage_scope
|
||||||
await db.execute(update(AIModelConfig).where(AIModelConfig.id != model_id).values(is_default=0))
|
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():
|
for field, val in req.model_dump(exclude_none=True).items():
|
||||||
if field == "api_key":
|
if field == "api_key":
|
||||||
model.api_key_enc = encrypt(val) if val else None
|
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="更新成功")
|
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}")
|
@router.delete("/{model_id}")
|
||||||
async def delete_model(model_id: int, db=Depends(get_db)):
|
async def delete_model(model_id: int, db=Depends(get_db)):
|
||||||
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
|
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:
|
def _format_model(m: AIModelConfig) -> dict:
|
||||||
return {
|
return {
|
||||||
"id": m.id, "model_name": m.model_name, "provider": m.provider,
|
"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),
|
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
|
||||||
"model_version": m.model_version, "temperature": m.temperature,
|
"model_version": m.model_version, "temperature": m.temperature,
|
||||||
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
|
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ class Settings(BaseSettings):
|
|||||||
# 安全
|
# 安全
|
||||||
SECRET_KEY: str = os.getenv("SECRET_KEY", "dev-secret-key-change-in-prod")
|
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!")
|
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(
|
NEWS_PLATFORM_BASE_URL: str = os.getenv(
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""数据库连接管理"""
|
"""数据库连接管理"""
|
||||||
import asyncio
|
import asyncio
|
||||||
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
||||||
|
from sqlalchemy import text
|
||||||
from sqlalchemy.orm import DeclarativeBase
|
from sqlalchemy.orm import DeclarativeBase
|
||||||
from app.core.config import settings
|
from app.core.config import settings
|
||||||
from app.core.logger import logger
|
from app.core.logger import logger
|
||||||
@@ -64,6 +65,18 @@ async def init_db():
|
|||||||
VirtualUser, UserPersonality, InteractionRecord,
|
VirtualUser, UserPersonality, InteractionRecord,
|
||||||
PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog
|
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("✅ 数据库模型注册成功")
|
||||||
logger.info("✅ 数据库初始化完成")
|
logger.info("✅ 数据库初始化完成")
|
||||||
|
|
||||||
|
|||||||
@@ -122,6 +122,7 @@ class AIModelConfig(Base):
|
|||||||
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
|
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
|
||||||
model_name: Mapped[str] = mapped_column(String(64), nullable=False)
|
model_name: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||||
provider: Mapped[str] = mapped_column(String(32), 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_base_url: Mapped[str | None] = mapped_column(String(256))
|
||||||
api_key_enc: Mapped[str | None] = mapped_column(String(512))
|
api_key_enc: Mapped[str | None] = mapped_column(String(512))
|
||||||
model_version: Mapped[str | None] = mapped_column(String(64))
|
model_version: Mapped[str | None] = mapped_column(String(64))
|
||||||
|
|||||||
@@ -154,6 +154,7 @@ class InteractionResponse(BaseModel):
|
|||||||
class AIModelCreateRequest(BaseModel):
|
class AIModelCreateRequest(BaseModel):
|
||||||
model_name: str = Field(..., min_length=1, max_length=64)
|
model_name: str = Field(..., min_length=1, max_length=64)
|
||||||
provider: str = Field(..., pattern="^(openai|zhipu|wenxin|qianwen|local)$")
|
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_base_url: Optional[str] = None
|
||||||
api_key: Optional[str] = None
|
api_key: Optional[str] = None
|
||||||
model_version: Optional[str] = None
|
model_version: Optional[str] = None
|
||||||
@@ -165,6 +166,8 @@ class AIModelCreateRequest(BaseModel):
|
|||||||
|
|
||||||
class AIModelUpdateRequest(BaseModel):
|
class AIModelUpdateRequest(BaseModel):
|
||||||
model_name: Optional[str] = None
|
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_base_url: Optional[str] = None
|
||||||
api_key: Optional[str] = None
|
api_key: Optional[str] = None
|
||||||
model_version: Optional[str] = None
|
model_version: Optional[str] = None
|
||||||
@@ -179,6 +182,7 @@ class AIModelResponse(BaseModel):
|
|||||||
id: int
|
id: int
|
||||||
model_name: str
|
model_name: str
|
||||||
provider: str
|
provider: str
|
||||||
|
usage_scope: str
|
||||||
api_base_url: Optional[str]
|
api_base_url: Optional[str]
|
||||||
has_api_key: bool
|
has_api_key: bool
|
||||||
model_version: Optional[str]
|
model_version: Optional[str]
|
||||||
|
|||||||
@@ -28,7 +28,9 @@ class AIService:
|
|||||||
async def _get_default_model(self, db: AsyncSession) -> Optional[AIModelConfig]:
|
async def _get_default_model(self, db: AsyncSession) -> Optional[AIModelConfig]:
|
||||||
result = await db.execute(
|
result = await db.execute(
|
||||||
select(AIModelConfig).where(
|
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()
|
return result.scalar_one_or_none()
|
||||||
|
|||||||
@@ -23,13 +23,10 @@ from services.token_billing import (
|
|||||||
reserve_avatar_tokens,
|
reserve_avatar_tokens,
|
||||||
settle_reservation,
|
settle_reservation,
|
||||||
)
|
)
|
||||||
|
from services.chat_model_config import ChatModelConfig, get_chat_model_config
|
||||||
|
|
||||||
router = APIRouter(tags=["数字分身聊天"])
|
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_MESSAGE_LENGTH = 4000
|
||||||
MAX_HISTORY_MESSAGES = 10
|
MAX_HISTORY_MESSAGES = 10
|
||||||
QA_LEXICAL_THRESHOLD = 0.72
|
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
|
return results
|
||||||
|
|
||||||
|
|
||||||
def _call_qwen(messages: list[dict], temperature: float) -> dict:
|
def _call_qwen(
|
||||||
if not CHAT_API_KEY:
|
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")
|
raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY")
|
||||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
url = f"{model_config.api_base_url}/chat/completions"
|
||||||
payload = {
|
payload = {
|
||||||
"model": CHAT_MODEL,
|
"model": model_config.model,
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
"max_tokens": model_config.max_tokens,
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
response = httpx.post(
|
response = httpx.post(
|
||||||
url,
|
url,
|
||||||
headers={"Authorization": f"Bearer {CHAT_API_KEY}"},
|
headers={"Authorization": f"Bearer {model_config.api_key}"},
|
||||||
json=payload,
|
json=payload,
|
||||||
timeout=30,
|
timeout=model_config.timeout_seconds,
|
||||||
)
|
)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
data = response.json()
|
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 {}}
|
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 分片原样转为文本增量。"""
|
"""将 OpenAI 兼容接口的 SSE 分片原样转为文本增量。"""
|
||||||
if not CHAT_API_KEY:
|
model_config = model_config or get_chat_model_config()
|
||||||
|
if not model_config.api_key:
|
||||||
raise RuntimeError("模型服务未配置")
|
raise RuntimeError("模型服务未配置")
|
||||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
url = f"{model_config.api_base_url}/chat/completions"
|
||||||
payload = {
|
payload = {
|
||||||
"model": CHAT_MODEL,
|
"model": model_config.model,
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
"max_tokens": model_config.max_tokens,
|
||||||
"stream": True,
|
"stream": True,
|
||||||
"stream_options": {"include_usage": True},
|
"stream_options": {"include_usage": True},
|
||||||
}
|
}
|
||||||
try:
|
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()
|
response.raise_for_status()
|
||||||
for raw_line in response.iter_lines():
|
for raw_line in response.iter_lines():
|
||||||
line = raw_line.decode() if isinstance(raw_line, bytes) else raw_line
|
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:
|
if model_client is not None:
|
||||||
answer = model_client(messages=messages, temperature=temperature)
|
answer = model_client(messages=messages, temperature=temperature)
|
||||||
else:
|
else:
|
||||||
|
model_config = get_chat_model_config()
|
||||||
reservation = reserve_avatar_tokens(
|
reservation = reserve_avatar_tokens(
|
||||||
db,
|
db,
|
||||||
avatar,
|
avatar,
|
||||||
usage_source,
|
usage_source,
|
||||||
CHAT_MODEL,
|
model_config.model,
|
||||||
messages,
|
messages,
|
||||||
CHAT_MAX_OUTPUT_TOKENS,
|
model_config.max_tokens,
|
||||||
)
|
)
|
||||||
try:
|
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"]
|
answer = model_result["answer"]
|
||||||
token_usage = settle_reservation(
|
token_usage = settle_reservation(
|
||||||
db,
|
db,
|
||||||
@@ -436,15 +450,16 @@ def _stream_reply(
|
|||||||
config = _config(avatar)
|
config = _config(avatar)
|
||||||
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||||
messages = _build_prompt(avatar, history, question, references)
|
messages = _build_prompt(avatar, history, question, references)
|
||||||
|
model_config = get_chat_model_config()
|
||||||
reservation = reserve_avatar_tokens(
|
reservation = reserve_avatar_tokens(
|
||||||
db,
|
db,
|
||||||
avatar,
|
avatar,
|
||||||
usage_source,
|
usage_source,
|
||||||
CHAT_MODEL,
|
model_config.model,
|
||||||
messages,
|
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:
|
if matched:
|
||||||
messages, reservation = [], None
|
messages, reservation = [], None
|
||||||
if public:
|
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"
|
||||||
@@ -10,6 +10,9 @@ services:
|
|||||||
environment:
|
environment:
|
||||||
DATABASE_URL: sqlite:////data/avatar.db
|
DATABASE_URL: sqlite:////data/avatar.db
|
||||||
UPLOAD_DIR: /data/uploads
|
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:
|
volumes:
|
||||||
- avatar-data:/data
|
- avatar-data:/data
|
||||||
expose:
|
expose:
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ services:
|
|||||||
- REDIS_PORT=6379
|
- REDIS_PORT=6379
|
||||||
- SECRET_KEY=your-secret-key-change-in-production
|
- SECRET_KEY=your-secret-key-change-in-production
|
||||||
- AES_KEY=your-aes-key-32-chars-change-now!
|
- AES_KEY=your-aes-key-32-chars-change-now!
|
||||||
|
- AVATAR_MODEL_CONFIG_TOKEN=${AVATAR_MODEL_CONFIG_TOKEN:-}
|
||||||
- TZ=Asia/Shanghai
|
- TZ=Asia/Shanghai
|
||||||
- AVATAR_DB_PATH=/app/avatar.db
|
- AVATAR_DB_PATH=/app/avatar.db
|
||||||
volumes:
|
volumes:
|
||||||
|
|||||||
@@ -94,6 +94,7 @@ CREATE TABLE IF NOT EXISTS `ai_model_configs` (
|
|||||||
`id` bigint NOT NULL AUTO_INCREMENT,
|
`id` bigint NOT NULL AUTO_INCREMENT,
|
||||||
`model_name` varchar(64) NOT NULL COMMENT '模型名称',
|
`model_name` varchar(64) NOT NULL COMMENT '模型名称',
|
||||||
`provider` varchar(32) NOT NULL COMMENT 'openai/zhipu/wenxin/qianwen/local',
|
`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_base_url` varchar(256) DEFAULT NULL COMMENT 'API地址',
|
||||||
`api_key_enc` varchar(512) DEFAULT NULL COMMENT '加密API Key',
|
`api_key_enc` varchar(512) DEFAULT NULL COMMENT '加密API Key',
|
||||||
`model_version` varchar(64) DEFAULT NULL COMMENT '模型版本',
|
`model_version` varchar(64) DEFAULT NULL COMMENT '模型版本',
|
||||||
|
|||||||
@@ -15,6 +15,9 @@
|
|||||||
<span class="model-title">{{ m.model_name }}</span>
|
<span class="model-title">{{ m.model_name }}</span>
|
||||||
</div>
|
</div>
|
||||||
<div style="display:flex;gap:6px;align-items:center">
|
<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_default" type="success" size="small">默认</el-tag>
|
||||||
<el-tag v-if="!m.is_enabled" type="danger" size="small">禁用</el-tag>
|
<el-tag v-if="!m.is_enabled" type="danger" size="small">禁用</el-tag>
|
||||||
</div>
|
</div>
|
||||||
@@ -49,6 +52,13 @@
|
|||||||
<el-option v-for="(l,v) in providerLabels" :key="v" :label="l" :value="v" />
|
<el-option v-for="(l,v) in providerLabels" :key="v" :label="l" :value="v" />
|
||||||
</el-select>
|
</el-select>
|
||||||
</el-form-item>
|
</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-form-item label="API地址">
|
||||||
<el-input v-model="form.api_base_url" placeholder="留空使用默认地址" />
|
<el-input v-model="form.api_base_url" placeholder="留空使用默认地址" />
|
||||||
</el-form-item>
|
</el-form-item>
|
||||||
@@ -131,8 +141,9 @@ const testResult = ref(null)
|
|||||||
const testing = ref(false)
|
const testing = ref(false)
|
||||||
|
|
||||||
const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' }
|
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 scopeLabels = { general: '通用业务', digital_avatar: '数字分身专用' }
|
||||||
const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }] }
|
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() {
|
async function load() {
|
||||||
const res = await getAIModels()
|
const res = await getAIModels()
|
||||||
@@ -155,13 +166,13 @@ function onProviderChange(provider) {
|
|||||||
|
|
||||||
function openCreate() {
|
function openCreate() {
|
||||||
editModel.value = null
|
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
|
dialogVisible.value = true
|
||||||
}
|
}
|
||||||
|
|
||||||
function openEdit(m) {
|
function openEdit(m) {
|
||||||
editModel.value = 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
|
dialogVisible.value = true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -236,4 +247,5 @@ onMounted(load)
|
|||||||
.result-meta { display: flex; align-items: center; gap: 10px; margin-bottom: 10px; }
|
.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; }
|
.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; }
|
.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>
|
</style>
|
||||||
|
|||||||
Reference in New Issue
Block a user