feat(ai): add dedicated digital avatar model config
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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("✅ 数据库初始化完成")
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user