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)