nex_docus/backend/app/api/v1/llm_model_configs.py

413 lines
13 KiB
Python

"""
LLM 模型配置 API
"""
from decimal import Decimal
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel, Field, field_validator
from sqlalchemy import func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.deps import get_current_user
from app.models.llm_model_config import LLMModelConfig
from app.models.user import User
from app.schemas.response import success_response
from app.services.llm_provider_service import LLMProviderService
router = APIRouter()
class LLMModelConfigUpsertRequest(BaseModel):
"""模型配置新增/编辑请求"""
model_code: Optional[str] = None
model_name: Optional[str] = None
provider: str = Field(..., min_length=1, max_length=64)
endpoint_url: Optional[str] = Field(None, max_length=512)
api_key: Optional[str] = Field(None, max_length=512)
llm_model_name: str = Field(..., min_length=1, max_length=128)
llm_timeout: int = Field(120, ge=5, le=600)
llm_temperature: float = Field(0.70, ge=0, le=2)
llm_top_p: float = Field(0.90, ge=0, le=1)
llm_max_tokens: int = Field(2048, ge=1, le=32768)
llm_system_prompt: Optional[str] = None
description: Optional[str] = Field(None, max_length=500)
is_active: bool = True
is_default: bool = False
@field_validator(
"model_code",
"model_name",
"provider",
"endpoint_url",
"api_key",
"llm_model_name",
"llm_system_prompt",
"description",
mode="before",
)
@classmethod
def strip_string_fields(cls, value):
if isinstance(value, str):
value = value.strip()
return value or None
return value
class LLMModelConfigTestRequest(LLMModelConfigUpsertRequest):
"""模型测试请求"""
def serialize_model_config(config: LLMModelConfig, include_api_key: bool = False) -> dict:
"""序列化模型配置"""
api_key = config.api_key or ""
data = {
"config_id": config.config_id,
"model_code": config.model_code,
"model_name": config.model_name,
"provider": config.provider,
"endpoint_url": config.endpoint_url,
"llm_model_name": config.llm_model_name,
"llm_timeout": config.llm_timeout,
"llm_temperature": float(config.llm_temperature or 0),
"llm_top_p": float(config.llm_top_p or 0),
"llm_max_tokens": config.llm_max_tokens,
"llm_system_prompt": config.llm_system_prompt,
"description": config.description,
"is_active": bool(config.is_active),
"is_default": bool(config.is_default),
"has_api_key": bool(api_key),
"api_key_masked": mask_api_key(api_key),
"created_at": config.created_at.isoformat() if config.created_at else None,
"updated_at": config.updated_at.isoformat() if config.updated_at else None,
}
if include_api_key:
data["api_key"] = api_key
return data
def mask_api_key(api_key: str) -> str:
"""脱敏 API Key"""
if not api_key:
return ""
if len(api_key) <= 8:
return "*" * len(api_key)
return f"{api_key[:4]}{'*' * (len(api_key) - 8)}{api_key[-4:]}"
def normalize_payload(payload: LLMModelConfigUpsertRequest) -> dict:
"""补齐自动生成字段"""
data = payload.model_dump()
provider = data["provider"]
llm_model_name = data["llm_model_name"]
data["model_name"] = data.get("model_name") or LLMProviderService.build_model_name(provider, llm_model_name)
data["model_code"] = data.get("model_code") or LLMProviderService.build_model_code(provider, llm_model_name)
data["endpoint_url"] = data.get("endpoint_url") or LLMProviderService.get_default_endpoint_url(provider)
if data["is_default"]:
data["is_active"] = True
data["llm_temperature"] = Decimal(str(data["llm_temperature"]))
data["llm_top_p"] = Decimal(str(data["llm_top_p"]))
return data
async def ensure_default_config(db: AsyncSession, preferred_config_id: Optional[int] = None):
"""确保始终存在一个默认启用模型"""
default_result = await db.execute(
select(LLMModelConfig.config_id).where(
LLMModelConfig.is_default == True,
LLMModelConfig.is_active == True,
)
)
if default_result.scalar_one_or_none():
return
candidate_id = None
if preferred_config_id:
candidate_result = await db.execute(
select(LLMModelConfig.config_id).where(
LLMModelConfig.config_id == preferred_config_id,
LLMModelConfig.is_active == True,
)
)
candidate_id = candidate_result.scalar_one_or_none()
if candidate_id is None:
fallback_result = await db.execute(
select(LLMModelConfig.config_id)
.where(LLMModelConfig.is_active == True)
.order_by(LLMModelConfig.updated_at.desc(), LLMModelConfig.config_id.desc())
.limit(1)
)
candidate_id = fallback_result.scalar_one_or_none()
if candidate_id is None:
return
await db.execute(update(LLMModelConfig).values(is_default=False))
await db.execute(
update(LLMModelConfig)
.where(LLMModelConfig.config_id == candidate_id)
.values(is_default=True)
)
@router.get("/providers", response_model=dict)
async def get_provider_catalog(
current_user: User = Depends(get_current_user),
):
"""获取模型提供方目录"""
return success_response(data=LLMProviderService.get_provider_catalog())
@router.get("/", response_model=dict)
async def get_llm_model_configs(
page: int = Query(1, ge=1),
page_size: int = Query(10, ge=1, le=100),
keyword: Optional[str] = Query(None, description="搜索关键词(模型名称、编码、模型名)"),
provider: Optional[str] = Query(None, description="提供方筛选"),
is_active: Optional[bool] = Query(None, description="启用状态筛选"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取模型配置列表"""
conditions = []
if keyword:
conditions.append(
or_(
LLMModelConfig.model_name.like(f"%{keyword}%"),
LLMModelConfig.model_code.like(f"%{keyword}%"),
LLMModelConfig.llm_model_name.like(f"%{keyword}%"),
)
)
if provider:
conditions.append(LLMModelConfig.provider == provider)
if is_active is not None:
conditions.append(LLMModelConfig.is_active == is_active)
count_query = select(func.count(LLMModelConfig.config_id))
if conditions:
count_query = count_query.where(*conditions)
total_result = await db.execute(count_query)
total = total_result.scalar() or 0
query = select(LLMModelConfig).order_by(
LLMModelConfig.is_default.desc(),
LLMModelConfig.updated_at.desc(),
LLMModelConfig.config_id.desc(),
)
if conditions:
query = query.where(*conditions)
query = query.offset((page - 1) * page_size).limit(page_size)
result = await db.execute(query)
configs = result.scalars().all()
return {
"code": 200,
"message": "success",
"data": [serialize_model_config(item) for item in configs],
"total": total,
"page": page,
"page_size": page_size,
}
@router.get("/{config_id}", response_model=dict)
async def get_llm_model_config_detail(
config_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取模型配置详情"""
result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.config_id == config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="模型配置不存在")
return success_response(data=serialize_model_config(config, include_api_key=True))
@router.post("/", response_model=dict)
async def create_llm_model_config(
request_data: LLMModelConfigUpsertRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""创建模型配置"""
payload = normalize_payload(request_data)
existing_code_result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.model_code == payload["model_code"])
)
if existing_code_result.scalar_one_or_none():
raise HTTPException(status_code=400, detail="模型编码已存在")
new_config = LLMModelConfig(**payload)
db.add(new_config)
await db.flush()
if payload["is_default"]:
await db.execute(
update(LLMModelConfig)
.where(LLMModelConfig.config_id != new_config.config_id)
.values(is_default=False)
)
await ensure_default_config(db, preferred_config_id=new_config.config_id)
await db.commit()
await db.refresh(new_config)
return success_response(
data=serialize_model_config(new_config),
message="模型配置创建成功",
)
@router.put("/{config_id}", response_model=dict)
async def update_llm_model_config(
config_id: int,
request_data: LLMModelConfigUpsertRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""更新模型配置"""
result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.config_id == config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="模型配置不存在")
payload = normalize_payload(request_data)
existing_code_result = await db.execute(
select(LLMModelConfig).where(
LLMModelConfig.model_code == payload["model_code"],
LLMModelConfig.config_id != config_id,
)
)
if existing_code_result.scalar_one_or_none():
raise HTTPException(status_code=400, detail="模型编码已被其他配置使用")
for key, value in payload.items():
setattr(config, key, value)
await db.flush()
if config.is_default:
await db.execute(
update(LLMModelConfig)
.where(LLMModelConfig.config_id != config.config_id)
.values(is_default=False)
)
await ensure_default_config(db, preferred_config_id=config.config_id)
await db.commit()
await db.refresh(config)
return success_response(
data=serialize_model_config(config),
message="模型配置更新成功",
)
@router.put("/{config_id}/status", response_model=dict)
async def update_llm_model_config_status(
config_id: int,
is_active: bool = Query(..., description="是否启用"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""更新模型配置启用状态"""
result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.config_id == config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="模型配置不存在")
config.is_active = is_active
if not is_active and config.is_default:
config.is_default = False
await db.flush()
await ensure_default_config(db)
await db.commit()
await db.refresh(config)
return success_response(
data=serialize_model_config(config),
message="模型状态更新成功",
)
@router.put("/{config_id}/default", response_model=dict)
async def set_default_llm_model_config(
config_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""设为默认模型配置"""
result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.config_id == config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="模型配置不存在")
config.is_active = True
config.is_default = True
await db.flush()
await db.execute(
update(LLMModelConfig)
.where(LLMModelConfig.config_id != config.config_id)
.values(is_default=False)
)
await db.commit()
await db.refresh(config)
return success_response(
data=serialize_model_config(config),
message="默认模型切换成功",
)
@router.delete("/{config_id}", response_model=dict)
async def delete_llm_model_config(
config_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""删除模型配置"""
result = await db.execute(
select(LLMModelConfig).where(LLMModelConfig.config_id == config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="模型配置不存在")
was_default = bool(config.is_default)
await db.delete(config)
await db.flush()
if was_default:
await ensure_default_config(db)
await db.commit()
return success_response(message="模型配置删除成功")
@router.post("/test", response_model=dict)
async def test_llm_model_config(
request_data: LLMModelConfigTestRequest,
current_user: User = Depends(get_current_user),
):
"""测试模型连接"""
payload = normalize_payload(request_data)
test_result = await LLMProviderService.test_model_connection(payload)
return success_response(data=test_result, message="模型测试成功")