Files
ai-agent-admin/backend/app/api/ai.py
T
2026-06-08 15:05:57 +08:00

171 lines
7.2 KiB
Python

import json
import time
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.crud import crud_router
from app.api.deps import get_current_user
from app.core.database import get_db
from app.models import (
Agent,
AgentConversation,
AgentMessage,
AgentTeam,
CollaborationRun,
KnowledgeBase,
LLMModel,
LLMProvider,
Workflow,
WorkflowRun,
)
from app.schemas.ai import (
AgentBase,
AgentOut,
ChatIn,
CollaborationRunIn,
CollaborationRunOut,
ConversationOut,
KnowledgeBaseIn,
KnowledgeBaseOut,
ModelBase,
ModelOut,
ProviderBase,
ProviderOut,
TeamBase,
TeamOut,
WorkflowBase,
WorkflowOut,
WorkflowRunIn,
WorkflowRunOut,
)
from app.schemas.common import Page, ResponseModel
from app.services.collaboration import CollaborationService
from app.services.llm import LLMService, sse
from app.services.workflow import WorkflowService
router = APIRouter(prefix="/ai", tags=["AI 平台"], dependencies=[Depends(get_current_user)])
router.include_router(
crud_router(prefix="/providers", tags=["Provider"], model=LLMProvider, create_schema=ProviderBase, update_schema=ProviderBase, out_schema=ProviderOut)
)
router.include_router(
crud_router(prefix="/models", tags=["Model"], model=LLMModel, create_schema=ModelBase, update_schema=ModelBase, out_schema=ModelOut)
)
router.include_router(
crud_router(prefix="/agents", tags=["Agent"], model=Agent, create_schema=AgentBase, update_schema=AgentBase, out_schema=AgentOut)
)
router.include_router(
crud_router(prefix="/workflows", tags=["Workflow"], model=Workflow, create_schema=WorkflowBase, update_schema=WorkflowBase, out_schema=WorkflowOut)
)
router.include_router(
crud_router(prefix="/knowledge-bases", tags=["Knowledge"], model=KnowledgeBase, create_schema=KnowledgeBaseIn, update_schema=KnowledgeBaseIn, out_schema=KnowledgeBaseOut)
)
router.include_router(
crud_router(prefix="/teams", tags=["Agent Team"], model=AgentTeam, create_schema=TeamBase, update_schema=TeamBase, out_schema=TeamOut)
)
@router.post("/agents/{agent_id}/chat")
async def chat(agent_id: str, payload: ChatIn, db: AsyncSession = Depends(get_db)):
agent = await db.get(Agent, agent_id)
if not agent or agent.is_deleted:
raise HTTPException(status_code=404, detail="智能体不存在")
conversation = None
if payload.conversation_id:
conversation = await db.get(AgentConversation, payload.conversation_id)
if not conversation:
conversation = AgentConversation(agent_id=agent.id, title=payload.message[:60] or "新建对话")
db.add(conversation)
await db.flush()
user_msg = AgentMessage(conversation_id=conversation.id, role="user", content=payload.message)
assistant_msg = AgentMessage(conversation_id=conversation.id, role="assistant", content="", status="pending")
db.add_all([user_msg, assistant_msg])
await db.flush()
async def generate():
started = time.time()
yield sse({"type": "start", "conversation_id": conversation.id, "message_id": assistant_msg.id})
chunks = []
service = LLMService(db)
async for chunk in service.stream(agent, payload.message):
chunks.append(chunk)
yield sse({"type": "chunk", "content": chunk})
assistant_msg.content = "".join(chunks)
assistant_msg.status = "completed"
assistant_msg.elapsed_time = int((time.time() - started) * 1000)
conversation.total_tokens = conversation.total_tokens + len(assistant_msg.content)
await db.commit()
yield sse({"type": "complete", "conversation_id": conversation.id, "message_id": assistant_msg.id})
yield sse("[DONE]")
return StreamingResponse(generate(), media_type="text/event-stream")
@router.get("/agents/{agent_id}/conversations", response_model=Page[ConversationOut])
async def list_conversations(agent_id: str, page: int = 1, page_size: int = 20, db: AsyncSession = Depends(get_db)):
query = select(AgentConversation).where(AgentConversation.agent_id == agent_id, AgentConversation.is_deleted == False)
result = await db.execute(query.order_by(AgentConversation.created_at.desc()).offset((page - 1) * page_size).limit(page_size))
total = len(result.scalars().all())
result = await db.execute(query.order_by(AgentConversation.created_at.desc()).offset((page - 1) * page_size).limit(page_size))
return Page(items=result.scalars().all(), total=total, page=page, page_size=page_size)
@router.get("/conversations/{conversation_id}/messages")
async def list_messages(conversation_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(
select(AgentMessage)
.where(AgentMessage.conversation_id == conversation_id, AgentMessage.is_deleted == False)
.order_by(AgentMessage.created_at.asc())
)
return result.scalars().all()
@router.post("/workflows/{workflow_id}/publish", response_model=WorkflowOut)
async def publish_workflow(workflow_id: str, db: AsyncSession = Depends(get_db)):
workflow = await db.get(Workflow, workflow_id)
if not workflow or workflow.is_deleted:
raise HTTPException(status_code=404, detail="工作流不存在")
workflow.status = "published"
workflow.published_version = workflow.version
workflow.published_definition = workflow.definition
await db.commit()
await db.refresh(workflow)
return workflow
@router.post("/workflows/{workflow_id}/run", response_model=WorkflowRunOut)
async def run_workflow(workflow_id: str, payload: WorkflowRunIn, db: AsyncSession = Depends(get_db)):
workflow = await db.get(Workflow, workflow_id)
if not workflow or workflow.is_deleted:
raise HTTPException(status_code=404, detail="工作流不存在")
return await WorkflowService(db).run(workflow, payload.inputs)
@router.get("/workflow-runs", response_model=Page[WorkflowRunOut])
async def list_workflow_runs(page: int = 1, page_size: int = 20, db: AsyncSession = Depends(get_db)):
query = select(WorkflowRun).where(WorkflowRun.is_deleted == False)
result = await db.execute(query.order_by(WorkflowRun.created_at.desc()).offset((page - 1) * page_size).limit(page_size))
items = result.scalars().all()
return Page(items=items, total=len(items), page=page, page_size=page_size)
@router.post("/teams/{team_id}/run", response_model=CollaborationRunOut)
async def run_team(team_id: str, payload: CollaborationRunIn, db: AsyncSession = Depends(get_db)):
team = await db.get(AgentTeam, team_id)
if not team or team.is_deleted:
raise HTTPException(status_code=404, detail="团队不存在")
return await CollaborationService(db).run(team, payload.task)
@router.get("/collaboration-runs", response_model=Page[CollaborationRunOut])
async def list_collaboration_runs(page: int = 1, page_size: int = 20, db: AsyncSession = Depends(get_db)):
query = select(CollaborationRun).where(CollaborationRun.is_deleted == False)
result = await db.execute(query.order_by(CollaborationRun.created_at.desc()).offset((page - 1) * page_size).limit(page_size))
items = result.scalars().all()
return Page(items=items, total=len(items), page=page, page_size=page_size)