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

523 lines
17 KiB
Python
Raw Normal View History

2026-04-08 17:03:57 +00:00
"""
LLM 模型配置 API
"""
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query
2026-08-03 05:59:52 +00:00
from pydantic import BaseModel, Field, field_validator, model_validator
2026-04-08 17:03:57 +00:00
from sqlalchemy import func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
2026-08-03 05:59:52 +00:00
from app.core.config import settings
2026-04-08 17:03:57 +00:00
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
2026-08-03 05:59:52 +00:00
from app.services.local_embedding_service import LocalEmbeddingService
2026-04-08 17:03:57 +00:00
router = APIRouter()
2026-08-03 05:59:52 +00:00
TYPE_SPECIFIC_FIELDS = {
"llm_temperature",
"llm_top_p",
"llm_max_tokens",
"llm_system_prompt",
"embedding_dimension",
"chunk_size",
"chunk_overlap",
}
2026-04-08 17:03:57 +00:00
class LLMModelConfigUpsertRequest(BaseModel):
"""模型配置新增/编辑请求"""
model_code: Optional[str] = None
model_name: Optional[str] = None
2026-07-24 07:37:22 +00:00
model_type: str = Field("chat", pattern="^(chat|embedding)$")
2026-04-08 17:03:57 +00:00
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)
2026-07-24 07:37:22 +00:00
llm_max_tokens: int = Field(8192, ge=1, le=32768)
2026-04-08 17:03:57 +00:00
llm_system_prompt: Optional[str] = None
2026-07-24 07:37:22 +00:00
embedding_dimension: Optional[int] = Field(None, ge=1, le=8192)
2026-08-03 05:59:52 +00:00
chunk_size: int = Field(settings.CHUNK_SIZE, ge=100, le=10000)
chunk_overlap: int = Field(settings.CHUNK_OVERLAP, ge=0, le=5000)
2026-04-08 17:03:57 +00:00
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
2026-08-03 05:59:52 +00:00
@model_validator(mode="after")
def validate_embedding_options(self):
if self.model_type == "embedding" and self.chunk_overlap >= self.chunk_size:
raise ValueError("分块重叠字符数必须小于分块字符数")
return self
2026-04-08 17:03:57 +00:00
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,
2026-07-24 07:37:22 +00:00
"model_type": config.model_type or "chat",
2026-04-08 17:03:57 +00:00
"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,
2026-07-24 07:37:22 +00:00
"embedding_dimension": config.embedding_dimension,
2026-08-03 05:59:52 +00:00
"chunk_size": config.chunk_size,
"chunk_overlap": config.chunk_overlap,
2026-04-08 17:03:57 +00:00
"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"]
2026-07-24 07:37:22 +00:00
model_type = data.get("model_type") or "chat"
data["model_type"] = model_type
2026-04-08 17:03:57 +00:00
data["model_name"] = data.get("model_name") or LLMProviderService.build_model_name(provider, llm_model_name)
2026-07-24 07:37:22 +00:00
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
2026-04-08 17:03:57 +00:00
data["endpoint_url"] = data.get("endpoint_url") or LLMProviderService.get_default_endpoint_url(provider)
if data["is_default"]:
data["is_active"] = True
2026-07-24 07:37:22 +00:00
if model_type == "embedding":
data["llm_system_prompt"] = None
2026-04-08 17:03:57 +00:00
return data
2026-08-03 05:59:52 +00:00
def to_storage_payload(payload: dict) -> dict:
"""将接口扁平字段归并为按模型类型区分的 JSON 参数。"""
data = dict(payload)
if data["model_type"] == "embedding":
type_config = {
"chunk_size": data["chunk_size"],
"chunk_overlap": data["chunk_overlap"],
}
if data.get("embedding_dimension") is not None:
type_config["dimension"] = data["embedding_dimension"]
else:
type_config = {
"temperature": data["llm_temperature"],
"top_p": data["llm_top_p"],
"max_tokens": data["llm_max_tokens"],
}
if data.get("llm_system_prompt"):
type_config["system_prompt"] = data["llm_system_prompt"]
for field in TYPE_SPECIFIC_FIELDS:
data.pop(field, None)
data["type_config"] = type_config
return data
async def ensure_default_config(
db: AsyncSession,
model_type: str,
preferred_config_id: Optional[int] = None,
):
"""确保指定模型类型始终存在一个默认启用配置。"""
2026-04-08 17:03:57 +00:00
default_result = await db.execute(
2026-08-03 05:59:52 +00:00
select(LLMModelConfig.config_id)
.where(
LLMModelConfig.model_type == model_type,
2026-04-08 17:03:57 +00:00
LLMModelConfig.is_default == True,
LLMModelConfig.is_active == True,
)
2026-08-03 05:59:52 +00:00
.limit(1)
2026-04-08 17:03:57 +00:00
)
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,
2026-08-03 05:59:52 +00:00
LLMModelConfig.model_type == model_type,
2026-04-08 17:03:57 +00:00
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)
2026-08-03 05:59:52 +00:00
.where(
LLMModelConfig.model_type == model_type,
LLMModelConfig.is_active == True,
)
2026-04-08 17:03:57 +00:00
.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
2026-08-03 05:59:52 +00:00
await db.execute(
update(LLMModelConfig)
.where(LLMModelConfig.model_type == model_type)
.values(is_default=False)
)
2026-04-08 17:03:57 +00:00
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),
):
"""获取模型提供方目录"""
2026-08-03 05:59:52 +00:00
catalog = [dict(item) for item in LLMProviderService.get_provider_catalog()]
for item in catalog:
item["default_chunk_size"] = settings.CHUNK_SIZE
item["default_chunk_overlap"] = settings.CHUNK_OVERLAP
if item["value"] == "local":
item["models"] = LocalEmbeddingService.list_available_models()
return success_response(data=catalog)
2026-04-08 17:03:57 +00:00
@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="提供方筛选"),
2026-07-24 07:37:22 +00:00
model_type: Optional[str] = Query(None, description="模型类型筛选: chat/embedding"),
2026-04-08 17:03:57 +00:00
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)
2026-07-24 07:37:22 +00:00
if model_type:
conditions.append(LLMModelConfig.model_type == model_type)
2026-04-08 17:03:57 +00:00
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),
):
"""创建模型配置"""
2026-08-03 05:59:52 +00:00
payload = to_storage_payload(normalize_payload(request_data))
2026-04-08 17:03:57 +00:00
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)
2026-08-03 05:59:52 +00:00
.where(
LLMModelConfig.model_type == new_config.model_type,
LLMModelConfig.config_id != new_config.config_id,
)
2026-04-08 17:03:57 +00:00
.values(is_default=False)
)
2026-08-03 05:59:52 +00:00
await ensure_default_config(
db,
new_config.model_type,
preferred_config_id=new_config.config_id,
)
2026-04-08 17:03:57 +00:00
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="模型配置不存在")
2026-08-03 05:59:52 +00:00
previous_model_type = config.model_type or "chat"
payload = to_storage_payload(normalize_payload(request_data))
2026-04-08 17:03:57 +00:00
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)
2026-08-03 05:59:52 +00:00
.where(
LLMModelConfig.model_type == config.model_type,
LLMModelConfig.config_id != config.config_id,
)
2026-04-08 17:03:57 +00:00
.values(is_default=False)
)
2026-08-03 05:59:52 +00:00
await ensure_default_config(
db,
config.model_type,
preferred_config_id=config.config_id,
)
if previous_model_type != config.model_type:
await ensure_default_config(db, previous_model_type)
2026-04-08 17:03:57 +00:00
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()
2026-08-03 05:59:52 +00:00
await ensure_default_config(db, config.model_type)
2026-04-08 17:03:57 +00:00
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)
2026-08-03 05:59:52 +00:00
.where(
LLMModelConfig.model_type == config.model_type,
LLMModelConfig.config_id != config.config_id,
)
2026-04-08 17:03:57 +00:00
.values(is_default=False)
)
await db.commit()
await db.refresh(config)
return success_response(
2026-08-03 05:59:52 +00:00
data={
**serialize_model_config(config),
"requires_revectorization": config.model_type == "embedding",
},
2026-04-08 17:03:57 +00:00
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)
2026-08-03 05:59:52 +00:00
model_type = config.model_type or "chat"
2026-04-08 17:03:57 +00:00
await db.delete(config)
await db.flush()
if was_default:
2026-08-03 05:59:52 +00:00
await ensure_default_config(db, model_type)
2026-04-08 17:03:57 +00:00
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),
):
2026-07-24 07:37:22 +00:00
"""测试模型连接按类型分流chat 走对话补全embedding 走向量接口)"""
2026-04-08 17:03:57 +00:00
payload = normalize_payload(request_data)
2026-07-24 07:37:22 +00:00
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
2026-04-08 17:03:57 +00:00
return success_response(data=test_result, message="模型测试成功")