86 lines
3.0 KiB
Python
86 lines
3.0 KiB
Python
"""平台/服务商业务逻辑"""
|
|
|
|
from typing import List, Optional
|
|
|
|
from fastapi import HTTPException, status
|
|
from sqlmodel import Session, select
|
|
|
|
from app.core import crypto
|
|
from app.models.provider import Provider, ProviderCategory
|
|
from app.schemas.provider import ProviderCreate, ProviderRead, ProviderUpdate
|
|
|
|
|
|
def _to_read(provider: Provider) -> ProviderRead:
|
|
read = ProviderRead.model_validate(provider)
|
|
read.has_api_config = bool(provider.api_config_encrypted)
|
|
return read
|
|
|
|
|
|
def list_providers(
|
|
session: Session,
|
|
category: Optional[ProviderCategory] = None,
|
|
enabled: Optional[bool] = None,
|
|
) -> List[ProviderRead]:
|
|
stmt = select(Provider)
|
|
if category is not None:
|
|
stmt = stmt.where(Provider.category == category)
|
|
if enabled is not None:
|
|
stmt = stmt.where(Provider.enabled == enabled)
|
|
stmt = stmt.order_by(Provider.category.asc(), Provider.slug.asc())
|
|
return [_to_read(p) for p in session.exec(stmt).all()]
|
|
|
|
|
|
def get_provider(session: Session, provider_id: int) -> ProviderRead:
|
|
provider = session.get(Provider, provider_id)
|
|
if not provider:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="平台不存在")
|
|
return _to_read(provider)
|
|
|
|
|
|
def create_provider(session: Session, data: ProviderCreate) -> ProviderRead:
|
|
exists = session.exec(select(Provider).where(Provider.slug == data.slug)).first()
|
|
if exists:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST, detail=f"slug 已存在:{data.slug}"
|
|
)
|
|
provider_data = data.model_dump(exclude={"api_config"})
|
|
provider = Provider(**provider_data)
|
|
provider.api_config_encrypted = crypto.encrypt(data.api_config)
|
|
session.add(provider)
|
|
session.commit()
|
|
session.refresh(provider)
|
|
return _to_read(provider)
|
|
|
|
|
|
def update_provider(
|
|
session: Session, provider_id: int, data: ProviderUpdate
|
|
) -> ProviderRead:
|
|
provider = session.get(Provider, provider_id)
|
|
if not provider:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="平台不存在")
|
|
fields = data.model_dump(exclude_unset=True, exclude={"api_config"})
|
|
for key, value in fields.items():
|
|
setattr(provider, key, value)
|
|
if data.api_config is not None:
|
|
provider.api_config_encrypted = crypto.encrypt(data.api_config)
|
|
session.add(provider)
|
|
session.commit()
|
|
session.refresh(provider)
|
|
return _to_read(provider)
|
|
|
|
|
|
def delete_provider(session: Session, provider_id: int) -> None:
|
|
provider = session.get(Provider, provider_id)
|
|
if not provider:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="平台不存在")
|
|
session.delete(provider)
|
|
session.commit()
|
|
|
|
|
|
def get_api_config_plain(session: Session, provider_id: int) -> Optional[str]:
|
|
"""获取平台 API 配置明文(受保护接口使用)"""
|
|
provider = session.get(Provider, provider_id)
|
|
if not provider:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="平台不存在")
|
|
return crypto.decrypt(provider.api_config_encrypted)
|