forked from xiaohei/taiji-AI-PAD
301 lines
8.9 KiB
Python
301 lines
8.9 KiB
Python
"""Tool catalogue endpoints."""
|
|
|
|
from typing import Optional, List
|
|
from uuid import UUID
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
|
from sqlalchemy import select, func, or_
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
import structlog
|
|
|
|
from database import get_db
|
|
from models import Tool as ToolModel
|
|
from schemas import ToolDefinition, ToolCreate, ToolUpdate, ToolResponse, PaginatedResponse
|
|
from app.auth import get_current_user
|
|
from app.resource_control import enforce_resource_control
|
|
from decimal import Decimal
|
|
|
|
logger = structlog.get_logger()
|
|
router = APIRouter(prefix="/tools", tags=["tools"])
|
|
|
|
|
|
@router.get("", response_model=PaginatedResponse[ToolResponse])
|
|
async def list_tools(
|
|
db: AsyncSession = Depends(get_db),
|
|
category: Optional[str] = Query(None, description="Filter by category"),
|
|
search: Optional[str] = Query(None, description="Search in name and description"),
|
|
is_active: Optional[bool] = Query(None, description="Filter by active status"),
|
|
is_public: Optional[bool] = Query(None, description="Filter by public visibility"),
|
|
page: int = Query(1, ge=1, description="Page number"),
|
|
page_size: int = Query(20, ge=1, le=100, description="Items per page"),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> PaginatedResponse[ToolResponse]:
|
|
"""List all available tools with filtering and pagination."""
|
|
|
|
# Build base query
|
|
query = select(ToolModel)
|
|
|
|
# Apply filters
|
|
filters = []
|
|
|
|
# User can see public tools or their own tools
|
|
if current_user.get("role") != "super_admin":
|
|
user_id_str = current_user.get("user_id")
|
|
try:
|
|
user_uuid = UUID(user_id_str) if user_id_str else None
|
|
except (ValueError, TypeError):
|
|
user_uuid = None
|
|
|
|
if user_uuid:
|
|
filters.append(
|
|
or_(
|
|
ToolModel.is_public == True,
|
|
ToolModel.owner_id == user_uuid
|
|
)
|
|
)
|
|
else:
|
|
# 如果用户ID无效,只显示公开工具
|
|
filters.append(ToolModel.is_public == True)
|
|
|
|
if category:
|
|
filters.append(ToolModel.category == category)
|
|
|
|
if search:
|
|
search_pattern = f"%{search}%"
|
|
filters.append(
|
|
or_(
|
|
ToolModel.name.ilike(search_pattern),
|
|
ToolModel.description.ilike(search_pattern)
|
|
)
|
|
)
|
|
|
|
if is_active is not None:
|
|
filters.append(ToolModel.is_active == is_active)
|
|
|
|
if is_public is not None:
|
|
filters.append(ToolModel.is_public == is_public)
|
|
|
|
if filters:
|
|
query = query.where(*filters)
|
|
|
|
# Get total count
|
|
count_query = select(func.count()).select_from(query.subquery())
|
|
total_result = await db.execute(count_query)
|
|
total = total_result.scalar_one()
|
|
|
|
# Apply pagination
|
|
query = query.offset((page - 1) * page_size).limit(page_size)
|
|
query = query.order_by(ToolModel.created_at.desc())
|
|
|
|
result = await db.execute(query)
|
|
tools = result.scalars().all()
|
|
|
|
logger.info(
|
|
"tools_listed",
|
|
user_id=current_user.get("user_id"),
|
|
total=total,
|
|
page=page,
|
|
filters={"category": category, "search": search}
|
|
)
|
|
|
|
return PaginatedResponse(
|
|
items=[ToolResponse.model_validate(tool) for tool in tools],
|
|
total=total,
|
|
page=page,
|
|
page_size=page_size,
|
|
pages=(total + page_size - 1) // page_size
|
|
)
|
|
|
|
|
|
@router.get("/{tool_id}", response_model=ToolResponse)
|
|
async def get_tool(
|
|
tool_id: UUID,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> ToolResponse:
|
|
"""Get a specific tool by ID."""
|
|
|
|
query = select(ToolModel).where(ToolModel.id == tool_id)
|
|
result = await db.execute(query)
|
|
tool = result.scalar_one_or_none()
|
|
|
|
if not tool:
|
|
raise HTTPException(status_code=404, detail="Tool not found")
|
|
|
|
# Check permissions
|
|
user_id_str = current_user.get("user_id")
|
|
try:
|
|
user_uuid = UUID(user_id_str) if user_id_str else None
|
|
except (ValueError, TypeError):
|
|
user_uuid = None
|
|
|
|
if (
|
|
not tool.is_public
|
|
and (user_uuid is None or tool.owner_id != user_uuid)
|
|
and current_user.get("role") != "super_admin"
|
|
):
|
|
raise HTTPException(status_code=403, detail="Access denied")
|
|
|
|
logger.info("tool_retrieved", tool_id=str(tool_id), user_id=user_id_str)
|
|
|
|
return ToolResponse.model_validate(tool)
|
|
|
|
|
|
@router.post("", response_model=ToolResponse, status_code=201)
|
|
async def create_tool(
|
|
tool_data: ToolCreate,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> ToolResponse:
|
|
"""Create a new tool."""
|
|
|
|
# 获取用户ID
|
|
user_id_str = current_user.get("user_id")
|
|
if not user_id_str:
|
|
raise HTTPException(status_code=401, detail="用户ID无效")
|
|
|
|
try:
|
|
user_uuid = UUID(user_id_str)
|
|
except (ValueError, TypeError):
|
|
raise HTTPException(status_code=401, detail="用户ID格式无效")
|
|
|
|
# ========== 资源管控检查 ==========
|
|
await enforce_resource_control(
|
|
user_id=user_id_str,
|
|
resource_type="tool",
|
|
resource_id=None,
|
|
estimated_cost=Decimal("0.0"), # 工具创建本身不收费
|
|
db=db
|
|
)
|
|
# ==================================
|
|
|
|
# Check if tool name already exists for this user
|
|
query = select(ToolModel).where(
|
|
ToolModel.name == tool_data.name,
|
|
ToolModel.owner_id == user_uuid
|
|
)
|
|
result = await db.execute(query)
|
|
existing = result.scalar_one_or_none()
|
|
|
|
if existing:
|
|
raise HTTPException(
|
|
status_code=409,
|
|
detail=f"Tool with name '{tool_data.name}' already exists"
|
|
)
|
|
|
|
# Create tool
|
|
tool = ToolModel(
|
|
**tool_data.model_dump(),
|
|
owner_id=user_uuid
|
|
)
|
|
|
|
db.add(tool)
|
|
await db.commit()
|
|
await db.refresh(tool)
|
|
|
|
logger.info(
|
|
"tool_created",
|
|
tool_id=str(tool.id),
|
|
tool_name=tool.name,
|
|
user_id=user_id_str
|
|
)
|
|
|
|
return ToolResponse.model_validate(tool)
|
|
|
|
|
|
@router.put("/{tool_id}", response_model=ToolResponse)
|
|
async def update_tool(
|
|
tool_id: UUID,
|
|
tool_data: ToolUpdate,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> ToolResponse:
|
|
"""Update an existing tool."""
|
|
|
|
query = select(ToolModel).where(ToolModel.id == tool_id)
|
|
result = await db.execute(query)
|
|
tool = result.scalar_one_or_none()
|
|
|
|
if not tool:
|
|
raise HTTPException(status_code=404, detail="Tool not found")
|
|
|
|
# 获取用户ID
|
|
user_id_str = current_user.get("user_id")
|
|
try:
|
|
user_uuid = UUID(user_id_str) if user_id_str else None
|
|
except (ValueError, TypeError):
|
|
user_uuid = None
|
|
|
|
# Check permissions
|
|
if (
|
|
(user_uuid is None or tool.owner_id != user_uuid)
|
|
and current_user.get("role") != "super_admin"
|
|
):
|
|
raise HTTPException(status_code=403, detail="Access denied")
|
|
|
|
# Update fields
|
|
update_data = tool_data.model_dump(exclude_unset=True)
|
|
for field, value in update_data.items():
|
|
setattr(tool, field, value)
|
|
|
|
await db.commit()
|
|
await db.refresh(tool)
|
|
|
|
logger.info(
|
|
"tool_updated",
|
|
tool_id=str(tool_id),
|
|
user_id=user_id_str,
|
|
updated_fields=list(update_data.keys())
|
|
)
|
|
|
|
return ToolResponse.model_validate(tool)
|
|
|
|
|
|
@router.delete("/{tool_id}", status_code=204)
|
|
async def delete_tool(
|
|
tool_id: UUID,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> None:
|
|
"""Delete a tool."""
|
|
|
|
query = select(ToolModel).where(ToolModel.id == tool_id)
|
|
result = await db.execute(query)
|
|
tool = result.scalar_one_or_none()
|
|
|
|
if not tool:
|
|
raise HTTPException(status_code=404, detail="Tool not found")
|
|
|
|
# 获取用户ID
|
|
user_id_str = current_user.get("user_id")
|
|
try:
|
|
user_uuid = UUID(user_id_str) if user_id_str else None
|
|
except (ValueError, TypeError):
|
|
user_uuid = None
|
|
|
|
# Check permissions
|
|
if (
|
|
(user_uuid is None or tool.owner_id != user_uuid)
|
|
and current_user.get("role") != "super_admin"
|
|
):
|
|
raise HTTPException(status_code=403, detail="Access denied")
|
|
|
|
await db.delete(tool)
|
|
await db.commit()
|
|
|
|
logger.info("tool_deleted", tool_id=str(tool_id), user_id=user_id_str)
|
|
|
|
|
|
@router.get("/categories/list", response_model=List[str])
|
|
async def list_categories(
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
) -> List[str]:
|
|
"""Get list of all tool categories."""
|
|
|
|
query = select(ToolModel.category).distinct().where(ToolModel.category.isnot(None))
|
|
result = await db.execute(query)
|
|
categories = [row[0] for row in result.all()]
|
|
|
|
return sorted(categories)
|