Files
taiji-AI-PAD/services/mcp-server/app/routes/tools.py
T

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)