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

435 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""
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
model_type: str = Field("chat", pattern="^(chat|embedding)$")
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(8192, ge=1, le=32768)
llm_system_prompt: Optional[str] = None
embedding_dimension: Optional[int] = Field(None, ge=1, le=8192)
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,
"model_type": config.model_type or "chat",
"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,
"embedding_dimension": config.embedding_dimension,
"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"]
model_type = data.get("model_type") or "chat"
data["model_type"] = model_type
data["model_name"] = data.get("model_name") or LLMProviderService.build_model_name(provider, llm_model_name)
base_code = data.get("model_code") or LLMProviderService.build_model_code(provider, llm_model_name)
# embedding 类型加前缀,避免与同名对话模型编码冲突
if model_type == "embedding" and not base_code.startswith("emb_"):
base_code = f"emb_{base_code}"
data["model_code"] = base_code
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"]))
# embedding 不需要对话采样参数,但列非空,保留默认值即可
if model_type == "embedding":
data["llm_system_prompt"] = None
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="提供方筛选"),
model_type: Optional[str] = Query(None, description="模型类型筛选: chat/embedding"),
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 model_type:
conditions.append(LLMModelConfig.model_type == model_type)
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),
):
"""测试模型连接按类型分流chat 走对话补全embedding 走向量接口)"""
payload = normalize_payload(request_data)
try:
if payload.get("model_type") == "embedding":
test_result = await LLMProviderService.test_embedding_connection(payload)
else:
test_result = await LLMProviderService.test_model_connection(payload)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return success_response(data=test_result, message="模型测试成功")