Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""
|
||||
AI 平台 API
|
||||
"""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,295 @@
|
||||
"""
|
||||
AI 应用 API
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.base_schema import PaginatedResponse, ResponseModel
|
||||
from ai_platform.models import AIApp, LLMModel
|
||||
from ai_platform.schemas.app_schema import (
|
||||
AppCreate,
|
||||
AppUpdate,
|
||||
AppResponse,
|
||||
AppListResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/apps", tags=["AI-应用"])
|
||||
|
||||
|
||||
@router.get("", response_model=PaginatedResponse[AppListResponse], summary="应用列表")
|
||||
async def list_apps(
|
||||
name: Optional[str] = Query(None, description="名称"),
|
||||
app_type: Optional[str] = Query(None, description="类型"),
|
||||
status: Optional[str] = Query(None, description="状态"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取应用列表"""
|
||||
query = select(AIApp).where(AIApp.is_deleted == False)
|
||||
|
||||
if name:
|
||||
query = query.where(AIApp.name.ilike(f"%{name}%"))
|
||||
if app_type:
|
||||
query = query.where(AIApp.app_type == app_type)
|
||||
if status:
|
||||
query = query.where(AIApp.status == status)
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(AIApp.sort.desc(), AIApp.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
apps = result.scalars().all()
|
||||
|
||||
items = []
|
||||
for app in apps:
|
||||
model_name = ""
|
||||
if app.model_id:
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(LLMModel.id == app.model_id)
|
||||
)
|
||||
model = model_result.scalar_one_or_none()
|
||||
model_name = model.display_name if model else ""
|
||||
|
||||
items.append({
|
||||
"id": app.id,
|
||||
"name": app.name,
|
||||
"code": app.code,
|
||||
"description": app.description or "",
|
||||
"icon": app.icon or "",
|
||||
"app_type": app.app_type or "chat",
|
||||
"status": app.status or "draft",
|
||||
"model_name": model_name,
|
||||
"is_public": app.is_public or False,
|
||||
"conversation_count": app.conversation_count or 0,
|
||||
"message_count": app.message_count or 0,
|
||||
"sys_create_datetime": app.sys_create_datetime,
|
||||
})
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.get("/published", response_model=List[AppListResponse], summary="获取已发布应用")
|
||||
async def list_published_apps(db: AsyncSession = Depends(get_db)):
|
||||
"""获取已发布的应用列表(用于用户选择)"""
|
||||
query = select(AIApp).where(
|
||||
AIApp.is_deleted == False,
|
||||
AIApp.status == "published"
|
||||
).order_by(AIApp.sort.desc())
|
||||
|
||||
result = await db.execute(query)
|
||||
apps = result.scalars().all()
|
||||
|
||||
items = []
|
||||
for app in apps:
|
||||
model_name = ""
|
||||
if app.model_id:
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(LLMModel.id == app.model_id)
|
||||
)
|
||||
model = model_result.scalar_one_or_none()
|
||||
model_name = model.display_name if model else ""
|
||||
|
||||
items.append({
|
||||
"id": app.id,
|
||||
"name": app.name,
|
||||
"code": app.code,
|
||||
"description": app.description or "",
|
||||
"icon": app.icon or "",
|
||||
"app_type": app.app_type or "chat",
|
||||
"status": app.status or "draft",
|
||||
"model_name": model_name,
|
||||
"is_public": app.is_public or False,
|
||||
"conversation_count": app.conversation_count or 0,
|
||||
"message_count": app.message_count or 0,
|
||||
"sys_create_datetime": app.sys_create_datetime,
|
||||
})
|
||||
|
||||
return items
|
||||
|
||||
|
||||
@router.get("/code/{code}", response_model=AppResponse, summary="根据编码获取应用")
|
||||
async def get_app_by_code(code: str, db: AsyncSession = Depends(get_db)):
|
||||
"""根据编码获取应用"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.code == code, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
@router.get("/{app_id}", response_model=AppResponse, summary="应用详情")
|
||||
async def get_app(app_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取应用详情"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
@router.post("", response_model=AppResponse, summary="创建应用")
|
||||
async def create_app(data: AppCreate, db: AsyncSession = Depends(get_db)):
|
||||
"""创建应用"""
|
||||
# 检查编码是否重复
|
||||
exists_result = await db.execute(
|
||||
select(AIApp).where(AIApp.code == data.code, AIApp.is_deleted == False)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail=f"应用编码 {data.code} 已存在")
|
||||
|
||||
# 验证模型
|
||||
if data.model_id:
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(LLMModel.id == data.model_id, LLMModel.is_deleted == False)
|
||||
)
|
||||
if not model_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail="模型不存在")
|
||||
|
||||
app = AIApp(**data.model_dump())
|
||||
db.add(app)
|
||||
await db.commit()
|
||||
await db.refresh(app)
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
@router.put("/{app_id}", response_model=AppResponse, summary="更新应用")
|
||||
async def update_app(
|
||||
app_id: str,
|
||||
data: AppUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新应用"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
|
||||
# 验证模型
|
||||
if "model_id" in update_data and update_data["model_id"]:
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == update_data["model_id"],
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
if not model_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail="模型不存在")
|
||||
|
||||
for key, value in update_data.items():
|
||||
setattr(app, key, value)
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(app)
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
@router.delete("/{app_id}", response_model=ResponseModel, summary="删除应用")
|
||||
async def delete_app(app_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""删除应用"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
app.is_deleted = True
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="删除成功")
|
||||
|
||||
|
||||
@router.post("/{app_id}/publish", response_model=AppResponse, summary="发布应用")
|
||||
async def publish_app(app_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""发布应用"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
app.status = "published"
|
||||
await db.commit()
|
||||
await db.refresh(app)
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
@router.post("/{app_id}/disable", response_model=AppResponse, summary="停用应用")
|
||||
async def disable_app(app_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""停用应用"""
|
||||
result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=404, detail="应用不存在")
|
||||
|
||||
app.status = "disabled"
|
||||
await db.commit()
|
||||
await db.refresh(app)
|
||||
|
||||
return await _build_app_response(app, db)
|
||||
|
||||
|
||||
async def _build_app_response(app: AIApp, db: AsyncSession) -> dict:
|
||||
"""构建应用输出"""
|
||||
model_name = ""
|
||||
if app.model_id:
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(LLMModel.id == app.model_id)
|
||||
)
|
||||
model = model_result.scalar_one_or_none()
|
||||
model_name = model.display_name if model else ""
|
||||
|
||||
return {
|
||||
"id": app.id,
|
||||
"name": app.name,
|
||||
"code": app.code,
|
||||
"description": app.description or "",
|
||||
"icon": app.icon or "",
|
||||
"app_type": app.app_type or "chat",
|
||||
"status": app.status or "draft",
|
||||
"model_id": app.model_id,
|
||||
"model_name": model_name,
|
||||
"system_prompt": app.system_prompt or "",
|
||||
"temperature": app.temperature or 0.7,
|
||||
"top_p": app.top_p or 1.0,
|
||||
"max_tokens": app.max_tokens or 2048,
|
||||
"opening_statement": app.opening_statement or "",
|
||||
"suggested_questions": app.suggested_questions or [],
|
||||
"workflow_definition": app.workflow_definition or {},
|
||||
"is_public": app.is_public or False,
|
||||
"conversation_count": app.conversation_count or 0,
|
||||
"message_count": app.message_count or 0,
|
||||
"sort": app.sort or 0,
|
||||
"sys_create_datetime": app.sys_create_datetime,
|
||||
"sys_update_datetime": app.sys_update_datetime,
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
"""
|
||||
对话 API
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.base_schema import PaginatedResponse, ResponseModel
|
||||
from utils.context import get_current_user_id_from_context
|
||||
from ai_platform.models import Conversation, Message, AIApp
|
||||
from ai_platform.schemas.chat_schema import (
|
||||
ConversationCreate,
|
||||
ConversationUpdate,
|
||||
ConversationResponse,
|
||||
ConversationListResponse,
|
||||
MessageResponse,
|
||||
SendMessageInput,
|
||||
MessageFeedbackInput,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/chat", tags=["AI-对话"])
|
||||
|
||||
|
||||
@router.get("/conversations", response_model=PaginatedResponse[ConversationListResponse], summary="对话列表")
|
||||
async def list_conversations(
|
||||
app_id: str = Query(..., description="应用 ID"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取对话列表"""
|
||||
query = select(Conversation).where(
|
||||
Conversation.app_id == app_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(
|
||||
Conversation.is_pinned.desc(),
|
||||
Conversation.sys_update_datetime.desc()
|
||||
)
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
conversations = result.scalars().all()
|
||||
|
||||
items = [
|
||||
{
|
||||
"id": c.id,
|
||||
"title": c.title or "",
|
||||
"message_count": c.message_count or 0,
|
||||
"is_pinned": c.is_pinned or False,
|
||||
"sys_update_datetime": c.sys_update_datetime,
|
||||
}
|
||||
for c in conversations
|
||||
]
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.post("/conversations", response_model=ConversationResponse, summary="创建对话")
|
||||
async def create_conversation(data: ConversationCreate, db: AsyncSession = Depends(get_db)):
|
||||
"""创建对话"""
|
||||
# 验证应用
|
||||
app_result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == data.app_id, AIApp.is_deleted == False)
|
||||
)
|
||||
app = app_result.scalar_one_or_none()
|
||||
if not app:
|
||||
raise HTTPException(status_code=400, detail="应用不存在")
|
||||
|
||||
user_id = get_current_user_id_from_context()
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=401, detail="未提供认证凭据")
|
||||
|
||||
conversation = Conversation(
|
||||
app_id=data.app_id,
|
||||
user_id=user_id,
|
||||
title=data.title or "新对话",
|
||||
)
|
||||
db.add(conversation)
|
||||
await db.commit()
|
||||
await db.refresh(conversation)
|
||||
|
||||
return await _build_conversation_response(conversation, db)
|
||||
|
||||
|
||||
@router.get("/conversations/{conversation_id}", response_model=ConversationResponse, summary="对话详情")
|
||||
async def get_conversation(conversation_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取对话详情"""
|
||||
result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
conversation = result.scalar_one_or_none()
|
||||
if not conversation:
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
return await _build_conversation_response(conversation, db)
|
||||
|
||||
|
||||
@router.put("/conversations/{conversation_id}", response_model=ConversationResponse, summary="更新对话")
|
||||
async def update_conversation(
|
||||
conversation_id: str,
|
||||
data: ConversationUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新对话"""
|
||||
result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
conversation = result.scalar_one_or_none()
|
||||
if not conversation:
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
if data.title is not None:
|
||||
conversation.title = data.title
|
||||
if data.is_pinned is not None:
|
||||
conversation.is_pinned = data.is_pinned
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(conversation)
|
||||
|
||||
return await _build_conversation_response(conversation, db)
|
||||
|
||||
|
||||
@router.delete("/conversations/{conversation_id}", response_model=ResponseModel, summary="删除对话")
|
||||
async def delete_conversation(conversation_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""删除对话"""
|
||||
result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
conversation = result.scalar_one_or_none()
|
||||
if not conversation:
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
conversation.is_deleted = True
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="删除成功")
|
||||
|
||||
|
||||
@router.get("/conversations/{conversation_id}/messages", response_model=List[MessageResponse], summary="获取消息列表")
|
||||
async def get_messages(
|
||||
conversation_id: str,
|
||||
limit: int = Query(50, ge=1, le=200, description="限制数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取对话消息"""
|
||||
# 验证对话存在
|
||||
conv_result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
if not conv_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
query = select(Message).where(
|
||||
Message.conversation_id == conversation_id,
|
||||
Message.is_deleted == False
|
||||
).order_by(Message.sys_create_datetime).limit(limit)
|
||||
|
||||
result = await db.execute(query)
|
||||
messages = result.scalars().all()
|
||||
|
||||
return [
|
||||
{
|
||||
"id": m.id,
|
||||
"role": m.role,
|
||||
"content": m.content or "",
|
||||
"status": m.status or "completed",
|
||||
"prompt_tokens": m.prompt_tokens or 0,
|
||||
"completion_tokens": m.completion_tokens or 0,
|
||||
"total_tokens": m.total_tokens or 0,
|
||||
"model_name": m.model_name or "",
|
||||
"latency": m.latency or 0,
|
||||
"error_message": m.error_message or "",
|
||||
"feedback": m.feedback or "",
|
||||
"sys_create_datetime": m.sys_create_datetime,
|
||||
}
|
||||
for m in messages
|
||||
]
|
||||
|
||||
|
||||
@router.post("/conversations/{conversation_id}/messages", response_model=MessageResponse, summary="发送消息")
|
||||
async def send_message(
|
||||
conversation_id: str,
|
||||
data: SendMessageInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""发送消息并获取 AI 回复"""
|
||||
from ai_platform.services.chat_service import ChatService
|
||||
|
||||
# 验证对话存在
|
||||
conv_result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
conversation = conv_result.scalar_one_or_none()
|
||||
if not conversation:
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
# 使用ChatService发送消息
|
||||
chat_service = ChatService(db)
|
||||
_, assistant_message = await chat_service.send_message(
|
||||
conversation_id=conversation_id,
|
||||
content=data.content,
|
||||
)
|
||||
|
||||
return {
|
||||
"id": assistant_message.id,
|
||||
"role": assistant_message.role,
|
||||
"content": assistant_message.content or "",
|
||||
"status": assistant_message.status or "completed",
|
||||
"prompt_tokens": assistant_message.prompt_tokens or 0,
|
||||
"completion_tokens": assistant_message.completion_tokens or 0,
|
||||
"total_tokens": assistant_message.total_tokens or 0,
|
||||
"model_name": assistant_message.model_name or "",
|
||||
"latency": assistant_message.latency or 0,
|
||||
"error_message": assistant_message.error_message or "",
|
||||
"feedback": assistant_message.feedback or "",
|
||||
"sys_create_datetime": assistant_message.sys_create_datetime,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/conversations/{conversation_id}/messages/stream", summary="流式发送消息")
|
||||
async def send_message_stream(
|
||||
conversation_id: str,
|
||||
data: SendMessageInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""发送消息并获取 AI 流式回复(SSE)"""
|
||||
from ai_platform.services.chat_service import ChatService
|
||||
|
||||
# 验证对话存在
|
||||
conv_result = await db.execute(
|
||||
select(Conversation).where(
|
||||
Conversation.id == conversation_id,
|
||||
Conversation.is_deleted == False
|
||||
)
|
||||
)
|
||||
conversation = conv_result.scalar_one_or_none()
|
||||
if not conversation:
|
||||
raise HTTPException(status_code=404, detail="对话不存在")
|
||||
|
||||
async def generate():
|
||||
chat_service = ChatService(db)
|
||||
async for event in chat_service.send_message_stream(
|
||||
conversation_id=conversation_id,
|
||||
content=data.content,
|
||||
):
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
generate(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.post("/messages/{message_id}/feedback", response_model=ResponseModel, summary="消息反馈")
|
||||
async def message_feedback(
|
||||
message_id: str,
|
||||
data: MessageFeedbackInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""对消息进行反馈"""
|
||||
result = await db.execute(
|
||||
select(Message).where(
|
||||
Message.id == message_id,
|
||||
Message.is_deleted == False
|
||||
)
|
||||
)
|
||||
message = result.scalar_one_or_none()
|
||||
if not message:
|
||||
raise HTTPException(status_code=404, detail="消息不存在")
|
||||
|
||||
if data.feedback not in ("like", "dislike", ""):
|
||||
raise HTTPException(status_code=400, detail="无效的反馈类型")
|
||||
|
||||
message.feedback = data.feedback
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="反馈成功")
|
||||
|
||||
|
||||
async def _build_conversation_response(conversation: Conversation, db: AsyncSession) -> dict:
|
||||
"""构建对话输出"""
|
||||
app_name = ""
|
||||
if conversation.app_id:
|
||||
app_result = await db.execute(
|
||||
select(AIApp).where(AIApp.id == conversation.app_id)
|
||||
)
|
||||
app = app_result.scalar_one_or_none()
|
||||
app_name = app.name if app else ""
|
||||
|
||||
return {
|
||||
"id": conversation.id,
|
||||
"app_id": conversation.app_id,
|
||||
"app_name": app_name,
|
||||
"title": conversation.title or "",
|
||||
"message_count": conversation.message_count or 0,
|
||||
"total_tokens": conversation.total_tokens or 0,
|
||||
"is_pinned": conversation.is_pinned or False,
|
||||
"sort": conversation.sort or 0,
|
||||
"sys_create_datetime": conversation.sys_create_datetime,
|
||||
"sys_update_datetime": conversation.sys_update_datetime,
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
"""
|
||||
LLM 模型 API
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Body
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.base_schema import PaginatedResponse, ResponseModel
|
||||
from ai_platform.models import LLMModel, LLMProvider
|
||||
from ai_platform.schemas.model_schema import (
|
||||
ModelCreate,
|
||||
ModelUpdate,
|
||||
ModelResponse,
|
||||
ModelListResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/model", tags=["AI-模型"])
|
||||
|
||||
|
||||
@router.get("/list", response_model=PaginatedResponse[ModelListResponse], summary="模型列表")
|
||||
async def list_models(
|
||||
provider_id: Optional[str] = Query(None, description="提供商 ID"),
|
||||
model_type: Optional[str] = Query(None, description="模型类型"),
|
||||
is_active: Optional[bool] = Query(None, description="是否启用"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取模型列表"""
|
||||
query = select(LLMModel).where(LLMModel.is_deleted == False)
|
||||
|
||||
if provider_id:
|
||||
query = query.where(LLMModel.provider_id == provider_id)
|
||||
if model_type:
|
||||
query = query.where(LLMModel.model_type == model_type)
|
||||
if is_active is not None:
|
||||
query = query.where(LLMModel.is_active == is_active)
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(LLMModel.sort.desc(), LLMModel.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
models = result.scalars().all()
|
||||
|
||||
# 获取提供商名称
|
||||
items = []
|
||||
for model in models:
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(LLMProvider.id == model.provider_id)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
items.append({
|
||||
"id": model.id,
|
||||
"provider_id": model.provider_id,
|
||||
"provider_name": provider.name if provider else "",
|
||||
"model_name": model.model_name,
|
||||
"display_name": model.display_name,
|
||||
"model_type": model.model_type or "chat",
|
||||
"is_active": model.is_active,
|
||||
})
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.get("/active", response_model=List[ModelListResponse], summary="获取可用模型列表")
|
||||
async def list_active_models(db: AsyncSession = Depends(get_db)):
|
||||
"""获取所有可用的模型(用于选择器)"""
|
||||
query = select(LLMModel).where(
|
||||
LLMModel.is_deleted == False,
|
||||
LLMModel.is_active == True
|
||||
).order_by(LLMModel.display_name)
|
||||
|
||||
result = await db.execute(query)
|
||||
models = result.scalars().all()
|
||||
|
||||
items = []
|
||||
for model in models:
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == model.provider_id,
|
||||
LLMProvider.is_deleted == False,
|
||||
LLMProvider.is_active == True
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if provider:
|
||||
items.append({
|
||||
"id": model.id,
|
||||
"provider_id": model.provider_id,
|
||||
"provider_name": provider.name,
|
||||
"model_name": model.model_name,
|
||||
"display_name": model.display_name,
|
||||
"model_type": model.model_type or "chat",
|
||||
"is_active": model.is_active,
|
||||
})
|
||||
|
||||
return items
|
||||
|
||||
|
||||
@router.get("/{model_id}", response_model=ModelResponse, summary="模型详情")
|
||||
async def get_model(model_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取模型详情"""
|
||||
result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise HTTPException(status_code=404, detail="模型不存在")
|
||||
|
||||
# 获取提供商名称
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(LLMProvider.id == model.provider_id)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
|
||||
return _build_model_response(model, provider)
|
||||
|
||||
|
||||
@router.post("", response_model=ModelResponse, summary="创建模型")
|
||||
async def create_model(data: ModelCreate, db: AsyncSession = Depends(get_db)):
|
||||
"""创建模型"""
|
||||
# 验证提供商
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == data.provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=400, detail="提供商不存在")
|
||||
|
||||
# 检查模型名称是否重复
|
||||
exists_result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.provider_id == data.provider_id,
|
||||
LLMModel.model_name == data.model_name,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail=f"模型 {data.model_name} 已存在")
|
||||
|
||||
model = LLMModel(**data.model_dump())
|
||||
db.add(model)
|
||||
await db.commit()
|
||||
await db.refresh(model)
|
||||
|
||||
return _build_model_response(model, provider)
|
||||
|
||||
|
||||
@router.post("/batch", response_model=ResponseModel, summary="批量创建模型")
|
||||
async def batch_create_models(
|
||||
provider_id: str = Query(..., description="提供商 ID"),
|
||||
models: List[dict] = Body(..., description="模型列表"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""批量创建模型(从默认模型列表)"""
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=400, detail="提供商不存在")
|
||||
|
||||
created_count = 0
|
||||
for model_data in models:
|
||||
model_name = model_data.get("model_name", "")
|
||||
if not model_name:
|
||||
continue
|
||||
|
||||
# 跳过已存在的
|
||||
exists_result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.provider_id == provider_id,
|
||||
LLMModel.model_name == model_name,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
continue
|
||||
|
||||
model = LLMModel(
|
||||
provider_id=provider_id,
|
||||
model_name=model_name,
|
||||
display_name=model_data.get("display_name", model_name),
|
||||
model_type=model_data.get("model_type", "chat"),
|
||||
max_tokens=model_data.get("max_tokens", 4096),
|
||||
context_window=model_data.get("context_window", 4096),
|
||||
supports_vision=model_data.get("supports_vision", False),
|
||||
supports_function_call=model_data.get("supports_function_call", False),
|
||||
input_price=model_data.get("input_price", 0),
|
||||
output_price=model_data.get("output_price", 0),
|
||||
is_active=True,
|
||||
)
|
||||
db.add(model)
|
||||
created_count += 1
|
||||
|
||||
await db.commit()
|
||||
return ResponseModel(message=f"成功创建 {created_count} 个模型")
|
||||
|
||||
|
||||
@router.put("/{model_id}", response_model=ModelResponse, summary="更新模型")
|
||||
async def update_model(
|
||||
model_id: str,
|
||||
data: ModelUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新模型"""
|
||||
result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise HTTPException(status_code=404, detail="模型不存在")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for key, value in update_data.items():
|
||||
setattr(model, key, value)
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(model)
|
||||
|
||||
# 获取提供商名称
|
||||
provider_result = await db.execute(
|
||||
select(LLMProvider).where(LLMProvider.id == model.provider_id)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
|
||||
return _build_model_response(model, provider)
|
||||
|
||||
|
||||
@router.delete("/{model_id}", response_model=ResponseModel, summary="删除模型")
|
||||
async def delete_model(model_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""删除模型"""
|
||||
result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise HTTPException(status_code=404, detail="模型不存在")
|
||||
|
||||
model.is_deleted = True
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="删除成功")
|
||||
|
||||
|
||||
def _build_model_response(model: LLMModel, provider: Optional[LLMProvider] = None) -> dict:
|
||||
"""构建模型输出"""
|
||||
return {
|
||||
"id": model.id,
|
||||
"provider_id": model.provider_id,
|
||||
"provider_name": provider.name if provider else "",
|
||||
"model_name": model.model_name,
|
||||
"display_name": model.display_name,
|
||||
"model_type": model.model_type or "chat",
|
||||
"max_tokens": model.max_tokens or 4096,
|
||||
"context_window": model.context_window or 4096,
|
||||
"default_temperature": model.default_temperature or 0.7,
|
||||
"default_top_p": model.default_top_p or 1.0,
|
||||
"input_price": model.input_price or 0,
|
||||
"output_price": model.output_price or 0,
|
||||
"supports_vision": model.supports_vision or False,
|
||||
"supports_function_call": model.supports_function_call or False,
|
||||
"supports_streaming": model.supports_streaming if model.supports_streaming is not None else True,
|
||||
"is_active": model.is_active,
|
||||
"sort": model.sort or 0,
|
||||
"sys_create_datetime": model.sys_create_datetime,
|
||||
"sys_update_datetime": model.sys_update_datetime,
|
||||
}
|
||||
@@ -0,0 +1,270 @@
|
||||
"""
|
||||
LLM 提供商 API
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.base_schema import PaginatedResponse, ResponseModel
|
||||
from ai_platform.models import LLMProvider
|
||||
from ai_platform.schemas.provider_schema import (
|
||||
ProviderCreate,
|
||||
ProviderUpdate,
|
||||
ProviderResponse,
|
||||
ProviderListResponse,
|
||||
ProviderTypeResponse,
|
||||
)
|
||||
from ai_platform.providers import ProviderRegistry
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/provider", tags=["AI-提供商"])
|
||||
|
||||
|
||||
@router.get("/types", response_model=List[ProviderTypeResponse], summary="获取提供商类型列表")
|
||||
async def get_provider_types():
|
||||
"""获取所有支持的提供商类型"""
|
||||
return ProviderRegistry.get_all_types()
|
||||
|
||||
|
||||
@router.get("/list", response_model=PaginatedResponse[ProviderListResponse], summary="提供商列表")
|
||||
async def list_providers(
|
||||
name: Optional[str] = Query(None, description="名称"),
|
||||
provider_type: Optional[str] = Query(None, description="类型"),
|
||||
is_active: Optional[bool] = Query(None, description="是否启用"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取提供商列表"""
|
||||
query = select(LLMProvider).where(LLMProvider.is_deleted == False)
|
||||
|
||||
if name:
|
||||
query = query.where(LLMProvider.name.ilike(f"%{name}%"))
|
||||
if provider_type:
|
||||
query = query.where(LLMProvider.provider_type == provider_type)
|
||||
if is_active is not None:
|
||||
query = query.where(LLMProvider.is_active == is_active)
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(LLMProvider.sort.desc(), LLMProvider.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.get("/{provider_id}", response_model=ProviderResponse, summary="提供商详情")
|
||||
async def get_provider(provider_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取提供商详情"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
return _build_provider_response(provider)
|
||||
|
||||
|
||||
@router.post("", response_model=ProviderResponse, summary="创建提供商")
|
||||
async def create_provider(data: ProviderCreate, db: AsyncSession = Depends(get_db)):
|
||||
"""创建提供商"""
|
||||
provider = LLMProvider(**data.model_dump())
|
||||
db.add(provider)
|
||||
await db.commit()
|
||||
await db.refresh(provider)
|
||||
|
||||
return _build_provider_response(provider)
|
||||
|
||||
|
||||
@router.put("/{provider_id}", response_model=ProviderResponse, summary="更新提供商")
|
||||
async def update_provider(
|
||||
provider_id: str,
|
||||
data: ProviderUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新提供商"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for key, value in update_data.items():
|
||||
setattr(provider, key, value)
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(provider)
|
||||
|
||||
return _build_provider_response(provider)
|
||||
|
||||
|
||||
@router.delete("/{provider_id}", response_model=ResponseModel, summary="删除提供商")
|
||||
async def delete_provider(provider_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""删除提供商"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
provider.is_deleted = True
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="删除成功")
|
||||
|
||||
|
||||
@router.post("/{provider_id}/test", summary="测试提供商连接")
|
||||
async def test_provider(provider_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""测试提供商连接"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
# 创建提供商实例
|
||||
provider_instance = ProviderRegistry.create_instance(
|
||||
provider_type=provider.provider_type,
|
||||
api_key=provider.api_key,
|
||||
api_base=provider.api_base,
|
||||
ollama_host=provider.ollama_host,
|
||||
)
|
||||
|
||||
if not provider_instance:
|
||||
raise HTTPException(status_code=400, detail=f"不支持的提供商类型: {provider.provider_type}")
|
||||
|
||||
# 验证配置
|
||||
if not provider_instance.validate_config():
|
||||
raise HTTPException(status_code=400, detail="配置无效,请检查 API Key")
|
||||
|
||||
# 尝试获取模型列表(如果支持)
|
||||
try:
|
||||
models = provider_instance.get_available_models()
|
||||
return {
|
||||
"success": True,
|
||||
"message": "连接成功",
|
||||
"models": models,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"获取模型列表失败: {e}")
|
||||
return {
|
||||
"success": True,
|
||||
"message": "连接成功(无法获取模型列表)",
|
||||
"models": [],
|
||||
}
|
||||
|
||||
|
||||
@router.get("/{provider_id}/default-models", summary="获取默认模型列表")
|
||||
async def get_default_models(provider_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取提供商的默认模型列表"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
return ProviderRegistry.get_default_models(provider.provider_type)
|
||||
|
||||
|
||||
@router.get("/{provider_id}/fetch-models", summary="在线拉取提供商模型列表")
|
||||
async def fetch_provider_models(provider_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""
|
||||
通过提供商 API 在线拉取最新模型列表。
|
||||
如果在线拉取失败,自动 fallback 到硬编码的默认模型列表。
|
||||
"""
|
||||
result = await db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == provider_id,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="提供商不存在")
|
||||
|
||||
# 构建 kwargs(Ollama 需要 ollama_host)
|
||||
kwargs = {}
|
||||
api_base = provider.api_base or ''
|
||||
if provider.provider_type == 'ollama' and provider.ollama_host:
|
||||
api_base = provider.ollama_host
|
||||
|
||||
# 尝试在线拉取
|
||||
try:
|
||||
models = await ProviderRegistry.fetch_models_from_api(
|
||||
provider_type=provider.provider_type,
|
||||
api_key=provider.api_key or '',
|
||||
api_base=api_base,
|
||||
**kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f'在线拉取模型列表异常 [provider={provider.name}, type={provider.provider_type}]: {e}'
|
||||
)
|
||||
models = []
|
||||
|
||||
source = 'api'
|
||||
if not models:
|
||||
logger.warning(
|
||||
f'在线拉取模型列表为空,fallback 到默认列表 '
|
||||
f'[provider={provider.name}, type={provider.provider_type}, api_base={api_base or "(empty)"}]'
|
||||
)
|
||||
models = ProviderRegistry.get_default_models(provider.provider_type)
|
||||
source = 'default'
|
||||
|
||||
return {
|
||||
"source": source,
|
||||
"models": models,
|
||||
}
|
||||
|
||||
|
||||
def _build_provider_response(provider: LLMProvider) -> dict:
|
||||
"""构建提供商输出"""
|
||||
return {
|
||||
"id": provider.id,
|
||||
"name": provider.name,
|
||||
"provider_type": provider.provider_type,
|
||||
"api_key_masked": provider.get_api_key_masked(),
|
||||
"api_base": provider.api_base or "",
|
||||
"api_version": provider.api_version or "",
|
||||
"ollama_host": provider.ollama_host or "",
|
||||
"description": provider.description or "",
|
||||
"is_active": provider.is_active,
|
||||
"quota_limit": provider.quota_limit or 0,
|
||||
"quota_used": provider.quota_used or 0,
|
||||
"sort": provider.sort or 0,
|
||||
"sys_create_datetime": provider.sys_create_datetime,
|
||||
"sys_update_datetime": provider.sys_update_datetime,
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
语音识别 API
|
||||
"""
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
||||
from fastapi.responses import Response
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/speech", tags=["AI-语音"])
|
||||
|
||||
|
||||
@router.post("/transcribe", summary="语音转文字")
|
||||
async def transcribe(
|
||||
audio: UploadFile = File(..., description="音频文件"),
|
||||
language: str = Query("zh", description="语言"),
|
||||
provider: str = Query("dashscope", description="提供商"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
语音转文字(ASR)
|
||||
|
||||
将音频文件转换为文字,支持:
|
||||
- dashscope: 阿里云百炼(默认,推荐)
|
||||
- openai: OpenAI Whisper
|
||||
|
||||
支持的音频格式:wav, mp3, webm, pcm, opus
|
||||
"""
|
||||
from ai_platform.services.speech_service import SpeechService
|
||||
|
||||
service = SpeechService(db=db)
|
||||
await service._resolve_dashscope_api_key()
|
||||
audio_content = await audio.read()
|
||||
result = service.transcribe(
|
||||
audio_file=audio_content,
|
||||
language=language,
|
||||
provider=provider,
|
||||
)
|
||||
|
||||
if result["success"]:
|
||||
return {
|
||||
"text": result["text"],
|
||||
"duration": result.get("duration", 0),
|
||||
}
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
|
||||
|
||||
@router.post("/tts", summary="文字转语音")
|
||||
async def text_to_speech(
|
||||
text: str = Query(..., description="要转换的文字"),
|
||||
voice: str = Query("sambert-zhichu-v1", description="声音"),
|
||||
provider: str = Query("dashscope", description="提供商"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
文字转语音(TTS)
|
||||
|
||||
将文字转换为语音,返回音频文件
|
||||
|
||||
DashScope 可用声音:
|
||||
- sambert-zhichu-v1: 知厨(男声)
|
||||
- sambert-zhimiao-emo-v1: 知妙(女声,带情感)
|
||||
- sambert-zhiying-v1: 知莺(女声)
|
||||
|
||||
OpenAI 可用声音:
|
||||
- alloy, echo, fable, onyx, nova, shimmer
|
||||
"""
|
||||
from ai_platform.services.speech_service import SpeechService
|
||||
|
||||
service = SpeechService(db=db)
|
||||
await service._resolve_dashscope_api_key()
|
||||
result = service.text_to_speech(
|
||||
text=text,
|
||||
voice=voice,
|
||||
provider=provider,
|
||||
)
|
||||
|
||||
if result["success"]:
|
||||
return Response(
|
||||
content=result["audio_data"],
|
||||
media_type=result["content_type"],
|
||||
headers={"Content-Disposition": 'attachment; filename="speech.wav"'},
|
||||
)
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
@@ -0,0 +1,720 @@
|
||||
"""
|
||||
AI 工作流 API
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, List, Optional
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import select, func, or_, and_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.database import get_db
|
||||
from app.base_schema import PaginatedResponse, ResponseModel
|
||||
from ai_platform.models import AIWorkflow, AIWorkflowVersion, AIWorkflowRun
|
||||
from core.application.model import Application
|
||||
from ai_platform.schemas.workflow_schema import (
|
||||
WorkflowCreate,
|
||||
WorkflowUpdate,
|
||||
WorkflowResponse,
|
||||
WorkflowListResponse,
|
||||
WorkflowRunInput,
|
||||
WorkflowRunResponse,
|
||||
WorkflowRunListResponse,
|
||||
WorkflowImportCheckIn,
|
||||
WorkflowImportCheckOut,
|
||||
WorkflowImportIn,
|
||||
NodeSchemaResponse,
|
||||
)
|
||||
from ai_platform.services.workflow_import_export import (
|
||||
WorkflowImportExportException,
|
||||
export_config as export_workflow_config,
|
||||
check_import as check_workflow_import,
|
||||
import_config as import_workflow_config,
|
||||
)
|
||||
from ai_platform.nodes.registry import NodeRegistry
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/workflow", tags=["AI-工作流"])
|
||||
|
||||
|
||||
# ============ 节点 Schema API ============
|
||||
|
||||
@router.get("/nodes/schemas", response_model=List[NodeSchemaResponse], summary="获取节点 Schema 列表")
|
||||
async def get_node_schemas():
|
||||
"""获取所有已注册节点的 Schema"""
|
||||
return NodeRegistry.get_all_schemas()
|
||||
|
||||
|
||||
@router.get("/nodes/schemas/by-category", summary="按分类获取节点 Schema")
|
||||
async def get_node_schemas_by_category():
|
||||
"""按分类获取所有已注册节点的 Schema"""
|
||||
return NodeRegistry.get_schemas_by_category()
|
||||
|
||||
|
||||
# ============ 工作流 API ============
|
||||
|
||||
@router.get("/list", response_model=PaginatedResponse[WorkflowListResponse], summary="工作流列表")
|
||||
async def list_workflows(
|
||||
name: Optional[str] = Query(None, description="名称"),
|
||||
status: Optional[str] = Query(None, description="状态"),
|
||||
workflow_type: Optional[str] = Query(None, description="工作流类型"),
|
||||
application_id: Optional[str] = Query(None, alias="applicationId", description="所属应用ID"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=1000, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取工作流列表(自动应用数据权限)"""
|
||||
from ai_platform.services.workflow_service import AIWorkflowService
|
||||
from app.data_scope_utils import get_data_scope_filter, apply_data_scope_to_conditions
|
||||
|
||||
conditions = [AIWorkflow.is_deleted == False]
|
||||
|
||||
if application_id:
|
||||
conditions.append(or_(
|
||||
AIWorkflow.application_id == application_id,
|
||||
and_(AIWorkflow.application_id.is_(None), AIWorkflow.is_global == True)
|
||||
))
|
||||
if name:
|
||||
conditions.append(AIWorkflow.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
conditions.append(AIWorkflow.status == status)
|
||||
if workflow_type:
|
||||
conditions.append(AIWorkflow.workflow_type == workflow_type)
|
||||
|
||||
# 获取数据权限过滤条件并应用
|
||||
from ai_platform.services.workflow_service import RESOURCE_TYPE
|
||||
data_scope_filter = await get_data_scope_filter(db, RESOURCE_TYPE)
|
||||
scope_conditions = apply_data_scope_to_conditions(AIWorkflow, data_scope_filter)
|
||||
conditions.extend(scope_conditions)
|
||||
|
||||
# 获取总数
|
||||
query = select(AIWorkflow).where(and_(*conditions))
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(AIWorkflow.sort.desc(), AIWorkflow.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
workflows = result.scalars().all()
|
||||
|
||||
# 批量查询应用名称
|
||||
app_ids = list({w.application_id for w in workflows if w.application_id})
|
||||
app_name_map = {}
|
||||
if app_ids:
|
||||
app_result = await db.execute(
|
||||
select(Application.id, Application.name).where(Application.id.in_(app_ids))
|
||||
)
|
||||
app_name_map = {row.id: row.name for row in app_result}
|
||||
|
||||
items = [_build_workflow_list_response(w, app_name_map.get(w.application_id, "")) for w in workflows]
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
# 全局运行记录路由须注册在 /{workflow_id} 之前,避免 "runs" 被当作 workflow_id
|
||||
@router.get("/runs", response_model=PaginatedResponse[WorkflowRunListResponse], summary="全局工作流运行记录")
|
||||
async def list_all_workflow_runs(
|
||||
workflow_id: Optional[str] = Query(None, alias="workflowId", description="工作流ID"),
|
||||
status: Optional[str] = Query(None, description="运行状态"),
|
||||
trigger_type: Optional[str] = Query(None, alias="triggerType", description="触发来源"),
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取全局工作流运行记录(可按工作流、状态、触发来源筛选)"""
|
||||
conditions = [AIWorkflowRun.is_deleted == False]
|
||||
if workflow_id:
|
||||
conditions.append(AIWorkflowRun.workflow_id == workflow_id)
|
||||
if status:
|
||||
conditions.append(AIWorkflowRun.status == status)
|
||||
if trigger_type:
|
||||
conditions.append(AIWorkflowRun.trigger_type == trigger_type)
|
||||
|
||||
query = select(AIWorkflowRun).where(*conditions).order_by(
|
||||
AIWorkflowRun.sys_create_datetime.desc()
|
||||
)
|
||||
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
result = await db.execute(query.offset(offset).limit(page_size))
|
||||
runs = result.scalars().all()
|
||||
|
||||
workflow_ids = {r.workflow_id for r in runs if r.workflow_id}
|
||||
workflow_name_map: dict[str, str] = {}
|
||||
if workflow_ids:
|
||||
wf_result = await db.execute(
|
||||
select(AIWorkflow.id, AIWorkflow.name).where(AIWorkflow.id.in_(workflow_ids))
|
||||
)
|
||||
workflow_name_map = {row[0]: row[1] for row in wf_result.all()}
|
||||
|
||||
items = [
|
||||
_build_run_list_item(r, workflow_name_map.get(r.workflow_id, ""))
|
||||
for r in runs
|
||||
]
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.get("/runs/{run_id}", response_model=WorkflowRunResponse, summary="运行记录详情")
|
||||
async def get_workflow_run(run_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取运行记录详情"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflowRun).where(AIWorkflowRun.id == run_id, AIWorkflowRun.is_deleted == False)
|
||||
)
|
||||
run = result.scalar_one_or_none()
|
||||
if not run:
|
||||
raise HTTPException(status_code=404, detail="运行记录不存在")
|
||||
|
||||
workflow_result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == run.workflow_id)
|
||||
)
|
||||
workflow = workflow_result.scalar_one_or_none()
|
||||
workflow_name = workflow.name if workflow else ""
|
||||
|
||||
return _build_run_detail(run, workflow_name)
|
||||
|
||||
|
||||
@router.post("/runs/{run_id}/stop", response_model=ResponseModel, summary="停止运行")
|
||||
async def stop_workflow_run(run_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""停止工作流运行"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflowRun).where(AIWorkflowRun.id == run_id, AIWorkflowRun.is_deleted == False)
|
||||
)
|
||||
run = result.scalar_one_or_none()
|
||||
if not run:
|
||||
raise HTTPException(status_code=404, detail="运行记录不存在")
|
||||
|
||||
if run.status not in ("pending", "running"):
|
||||
raise HTTPException(status_code=400, detail="工作流已完成,无法停止")
|
||||
|
||||
run.status = "stopped"
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="已停止")
|
||||
|
||||
|
||||
class ResumeWorkflowInput(BaseModel):
|
||||
"""恢复工作流输入"""
|
||||
user_input: Any = Field(..., description="用户输入(可以是字符串、布尔值、对象等)")
|
||||
|
||||
|
||||
@router.post("/runs/{run_id}/resume", summary="恢复工作流运行")
|
||||
async def resume_workflow_run(
|
||||
run_id: str,
|
||||
data: ResumeWorkflowInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""恢复等待中的工作流运行(SSE)"""
|
||||
import json
|
||||
from ai_platform.services.workflow_service import AIWorkflowService
|
||||
|
||||
async def generate():
|
||||
workflow_service = AIWorkflowService(db)
|
||||
async for event in workflow_service.resume_workflow_stream_async(
|
||||
run_id=run_id,
|
||||
user_input=data.user_input,
|
||||
):
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False, default=str)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
generate(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.post("/runs/{run_id}/resume/stream", summary="流式恢复工作流运行")
|
||||
async def resume_workflow_run_stream(
|
||||
run_id: str,
|
||||
data: ResumeWorkflowInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""流式恢复等待中的工作流运行(SSE)"""
|
||||
import json
|
||||
from ai_platform.services.workflow_service import AIWorkflowService
|
||||
|
||||
async def generate():
|
||||
workflow_service = AIWorkflowService(db)
|
||||
async for event in workflow_service.resume_workflow_stream_async(
|
||||
run_id=run_id,
|
||||
user_input=data.user_input,
|
||||
):
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False, default=str)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
generate(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/code/{code}", response_model=WorkflowResponse, summary="根据编码获取工作流")
|
||||
async def get_workflow_by_code(code: str, db: AsyncSession = Depends(get_db)):
|
||||
"""根据编码获取工作流"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.code == code, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
return _build_workflow_response(workflow)
|
||||
|
||||
|
||||
@router.get("/{workflow_id}", response_model=WorkflowResponse, summary="工作流详情")
|
||||
async def get_workflow(workflow_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""获取工作流详情"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
return _build_workflow_response(workflow)
|
||||
|
||||
|
||||
@router.post("", response_model=WorkflowResponse, summary="创建工作流")
|
||||
async def create_workflow(data: WorkflowCreate, db: AsyncSession = Depends(get_db)):
|
||||
"""创建工作流"""
|
||||
# 检查编码是否重复
|
||||
exists_result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.code == data.code, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail=f"工作流编码 {data.code} 已存在")
|
||||
|
||||
workflow = AIWorkflow(**data.model_dump())
|
||||
db.add(workflow)
|
||||
await db.commit()
|
||||
await db.refresh(workflow)
|
||||
|
||||
return _build_workflow_response(workflow)
|
||||
|
||||
|
||||
@router.put("/{workflow_id}", response_model=WorkflowResponse, summary="更新工作流")
|
||||
async def update_workflow(
|
||||
workflow_id: str,
|
||||
data: WorkflowUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""更新工作流"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
|
||||
# 如果更新了 code,检查编码是否重复(排除自身)
|
||||
if data.code and data.code != workflow.code:
|
||||
exists_result = await db.execute(
|
||||
select(AIWorkflow).where(
|
||||
AIWorkflow.code == data.code,
|
||||
AIWorkflow.id != workflow_id,
|
||||
AIWorkflow.is_deleted == False
|
||||
)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail=f"工作流编码 {data.code} 已存在")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for key, value in update_data.items():
|
||||
setattr(workflow, key, value)
|
||||
|
||||
# 更新草稿版本号
|
||||
workflow.version = (workflow.version or 0) + 1
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(workflow)
|
||||
|
||||
return _build_workflow_response(workflow)
|
||||
|
||||
|
||||
@router.delete("/{workflow_id}", response_model=ResponseModel, summary="删除工作流")
|
||||
async def delete_workflow(workflow_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""删除工作流"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
|
||||
workflow.is_deleted = True
|
||||
await db.commit()
|
||||
|
||||
return ResponseModel(message="删除成功")
|
||||
|
||||
|
||||
@router.post("/{workflow_id}/copy", response_model=WorkflowResponse, summary="复制工作流")
|
||||
async def copy_workflow(workflow_id: str, db: AsyncSession = Depends(get_db)):
|
||||
"""复制工作流"""
|
||||
# 获取原工作流
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
|
||||
# 生成新的编码(原编码 + _copy + 时间戳)
|
||||
import time
|
||||
timestamp = int(time.time() * 1000)
|
||||
new_code = f"{workflow.code}_copy_{timestamp}"
|
||||
|
||||
# 检查编码是否重复(理论上不会,但保险起见)
|
||||
exists_result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.code == new_code, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
if exists_result.scalar_one_or_none():
|
||||
new_code = f"{workflow.code}_copy_{timestamp}_{int(time.time())}"
|
||||
|
||||
# 创建新工作流
|
||||
new_workflow = AIWorkflow(
|
||||
name=f"{workflow.name} (副本)",
|
||||
code=new_code,
|
||||
description=workflow.description,
|
||||
workflow_type=workflow.workflow_type,
|
||||
definition=workflow.definition, # 复制工作流定义
|
||||
input_variables=workflow.input_variables,
|
||||
output_variables=workflow.output_variables,
|
||||
status="draft", # 新工作流默认为草稿状态
|
||||
version=1,
|
||||
published_version=None,
|
||||
published_at=None,
|
||||
published_definition=None,
|
||||
)
|
||||
db.add(new_workflow)
|
||||
await db.commit()
|
||||
await db.refresh(new_workflow)
|
||||
|
||||
return _build_workflow_response(new_workflow)
|
||||
|
||||
|
||||
@router.get("/{workflow_id}/export", summary="导出工作流配置")
|
||||
async def export_workflow(
|
||||
workflow_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""导出工作流配置为 JSON(草稿 definition)"""
|
||||
try:
|
||||
config = await export_workflow_config(db, workflow_id)
|
||||
content = json.dumps(config, ensure_ascii=False, indent=2)
|
||||
return StreamingResponse(
|
||||
iter([content]),
|
||||
media_type="application/json",
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="{config["code"]}.json"'
|
||||
},
|
||||
)
|
||||
except WorkflowImportExportException as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
|
||||
@router.post("/import/check", response_model=WorkflowImportCheckOut, summary="导入预检查")
|
||||
async def check_import_workflow(
|
||||
data: WorkflowImportCheckIn,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""导入预检查:检查工作流编码是否冲突"""
|
||||
try:
|
||||
return await check_workflow_import(db, data.code)
|
||||
except WorkflowImportExportException as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
|
||||
@router.post("/import", response_model=WorkflowResponse, summary="导入工作流配置")
|
||||
async def import_workflow(
|
||||
data: WorkflowImportIn,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""导入工作流配置"""
|
||||
try:
|
||||
workflow = await import_workflow_config(db, data.model_dump())
|
||||
await db.commit()
|
||||
await db.refresh(workflow)
|
||||
return _build_workflow_response(workflow)
|
||||
except WorkflowImportExportException as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
|
||||
@router.post("/{workflow_id}/publish", response_model=WorkflowResponse, summary="发布工作流")
|
||||
async def publish_workflow(
|
||||
workflow_id: str,
|
||||
description: str = Query("", description="版本说明"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""发布工作流"""
|
||||
result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||||
)
|
||||
workflow = result.scalar_one_or_none()
|
||||
if not workflow:
|
||||
raise HTTPException(status_code=404, detail="工作流不存在")
|
||||
|
||||
# 创建版本记录
|
||||
new_version = (workflow.published_version or 0) + 1
|
||||
version = AIWorkflowVersion(
|
||||
workflow_id=workflow_id,
|
||||
version=new_version,
|
||||
definition=workflow.definition or {},
|
||||
description=description,
|
||||
published_at=datetime.now(),
|
||||
)
|
||||
db.add(version)
|
||||
|
||||
# 更新工作流
|
||||
workflow.status = "published"
|
||||
workflow.published_version = new_version
|
||||
workflow.published_at = datetime.now()
|
||||
workflow.published_definition = workflow.definition
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(workflow)
|
||||
|
||||
return _build_workflow_response(workflow)
|
||||
|
||||
|
||||
@router.get("/{workflow_id}/versions", summary="获取版本历史")
|
||||
async def list_workflow_versions(
|
||||
workflow_id: str,
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取工作流版本历史"""
|
||||
query = select(AIWorkflowVersion).where(
|
||||
AIWorkflowVersion.workflow_id == workflow_id,
|
||||
AIWorkflowVersion.is_deleted == False
|
||||
).order_by(AIWorkflowVersion.version.desc())
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
versions = result.scalars().all()
|
||||
|
||||
items = [
|
||||
{
|
||||
"id": v.id,
|
||||
"version": v.version,
|
||||
"description": v.description or "",
|
||||
"published_at": v.published_at,
|
||||
"run_count": v.run_count or 0,
|
||||
"success_count": v.success_count or 0,
|
||||
}
|
||||
for v in versions
|
||||
]
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
@router.get("/{workflow_id}/runs", response_model=PaginatedResponse[WorkflowRunListResponse], summary="工作流运行记录")
|
||||
async def list_workflow_runs(
|
||||
workflow_id: str,
|
||||
page: int = Query(1, ge=1, description="页码"),
|
||||
page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取工作流运行记录"""
|
||||
query = select(AIWorkflowRun).where(
|
||||
AIWorkflowRun.workflow_id == workflow_id,
|
||||
AIWorkflowRun.is_deleted == False
|
||||
).order_by(AIWorkflowRun.sys_create_datetime.desc())
|
||||
|
||||
# 获取总数
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await db.execute(query)
|
||||
runs = result.scalars().all()
|
||||
|
||||
# 获取工作流名称
|
||||
workflow_result = await db.execute(
|
||||
select(AIWorkflow).where(AIWorkflow.id == workflow_id)
|
||||
)
|
||||
workflow = workflow_result.scalar_one_or_none()
|
||||
workflow_name = workflow.name if workflow else ""
|
||||
|
||||
items = [
|
||||
_build_run_list_item(r, workflow_name)
|
||||
for r in runs
|
||||
]
|
||||
|
||||
return PaginatedResponse(items=items, total=total)
|
||||
|
||||
|
||||
def _build_run_list_item(run: AIWorkflowRun, workflow_name: str = "") -> dict:
|
||||
return {
|
||||
"id": run.id,
|
||||
"workflow_id": run.workflow_id,
|
||||
"workflow_name": workflow_name,
|
||||
"status": run.status or "pending",
|
||||
"trigger_type": run.trigger_type or "api",
|
||||
"total_steps": run.total_steps or 0,
|
||||
"total_tokens": run.total_tokens or 0,
|
||||
"elapsed_time": run.elapsed_time or 0,
|
||||
"error_message": (run.error_message or "")[:200],
|
||||
"started_at": run.started_at,
|
||||
"completed_at": run.completed_at,
|
||||
}
|
||||
|
||||
|
||||
def _build_run_detail(run: AIWorkflowRun, workflow_name: str = "") -> dict:
|
||||
return {
|
||||
"id": run.id,
|
||||
"workflow_id": run.workflow_id,
|
||||
"workflow_name": workflow_name,
|
||||
"status": run.status or "pending",
|
||||
"trigger_type": run.trigger_type or "api",
|
||||
"use_draft": bool(run.use_draft),
|
||||
"workflow_version": run.workflow_version,
|
||||
"definition_snapshot": run.definition_snapshot or {},
|
||||
"inputs": run.inputs or {},
|
||||
"outputs": run.outputs or {},
|
||||
"execution_log": run.execution_log or [],
|
||||
"current_node_id": run.current_node_id or "",
|
||||
"waiting_config": run.waiting_config or {},
|
||||
"error_message": run.error_message or "",
|
||||
"total_tokens": run.total_tokens or 0,
|
||||
"total_steps": run.total_steps or 0,
|
||||
"elapsed_time": run.elapsed_time or 0,
|
||||
"started_at": run.started_at,
|
||||
"completed_at": run.completed_at,
|
||||
"sys_create_datetime": run.sys_create_datetime,
|
||||
}
|
||||
|
||||
|
||||
class WorkflowRunInput(BaseModel):
|
||||
"""工作流运行输入"""
|
||||
inputs: dict = Field(default_factory=dict, description="输入变量")
|
||||
use_draft: bool = Field(default=False, description="是否使用草稿版本")
|
||||
|
||||
|
||||
@router.post("/{workflow_id}/run", summary="运行工作流")
|
||||
async def run_workflow(
|
||||
workflow_id: str,
|
||||
data: WorkflowRunInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""运行工作流(非流式)"""
|
||||
from ai_platform.services.workflow_service import AIWorkflowService
|
||||
|
||||
workflow_service = AIWorkflowService(db)
|
||||
run = await workflow_service.run_workflow(
|
||||
workflow_id=workflow_id,
|
||||
inputs=data.inputs,
|
||||
use_draft=data.use_draft,
|
||||
trigger_type='api',
|
||||
)
|
||||
|
||||
return {
|
||||
"id": run.id,
|
||||
"workflow_id": run.workflow_id,
|
||||
"status": run.status,
|
||||
"outputs": run.outputs or {},
|
||||
"execution_log": run.execution_log or [],
|
||||
"total_tokens": run.total_tokens or 0,
|
||||
"total_steps": run.total_steps or 0,
|
||||
"elapsed_time": run.elapsed_time or 0,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/{workflow_id}/run/stream", summary="流式运行工作流")
|
||||
async def run_workflow_stream(
|
||||
workflow_id: str,
|
||||
data: WorkflowRunInput,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""流式运行工作流(SSE)"""
|
||||
import json
|
||||
from ai_platform.services.workflow_service import AIWorkflowService
|
||||
|
||||
async def generate():
|
||||
workflow_service = AIWorkflowService(db)
|
||||
trigger_type = 'editor_draft' if data.use_draft else 'editor_published'
|
||||
async for event in workflow_service.run_workflow_stream_async(
|
||||
workflow_id=workflow_id,
|
||||
inputs=data.inputs,
|
||||
use_draft=data.use_draft,
|
||||
trigger_type=trigger_type,
|
||||
):
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False, default=str)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
generate(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _build_workflow_response(workflow: AIWorkflow) -> dict:
|
||||
"""构建工作流输出"""
|
||||
return {
|
||||
"id": workflow.id,
|
||||
"name": workflow.name,
|
||||
"code": workflow.code,
|
||||
"workflow_type": workflow.workflow_type or "general",
|
||||
"description": workflow.description or "",
|
||||
"status": workflow.status or "draft",
|
||||
"version": workflow.version or 1,
|
||||
"published_version": workflow.published_version,
|
||||
"published_at": workflow.published_at,
|
||||
"published_definition": workflow.published_definition,
|
||||
"definition": workflow.definition or {},
|
||||
"input_variables": workflow.input_variables or [],
|
||||
"output_variables": workflow.output_variables or [],
|
||||
"run_count": workflow.run_count or 0,
|
||||
"success_count": workflow.success_count or 0,
|
||||
"sort": workflow.sort or 0,
|
||||
"sys_create_datetime": workflow.sys_create_datetime,
|
||||
"sys_update_datetime": workflow.sys_update_datetime,
|
||||
}
|
||||
|
||||
|
||||
def _build_workflow_list_response(workflow: AIWorkflow, application_name: str = "") -> dict:
|
||||
"""构建工作流列表输出"""
|
||||
return {
|
||||
"id": workflow.id,
|
||||
"application_id": workflow.application_id,
|
||||
"application_name": application_name,
|
||||
"is_global": workflow.is_global or False,
|
||||
"name": workflow.name,
|
||||
"code": workflow.code,
|
||||
"workflow_type": workflow.workflow_type or "general",
|
||||
"description": workflow.description or "",
|
||||
"status": workflow.status or "draft",
|
||||
"version": workflow.version or 1,
|
||||
"run_count": workflow.run_count or 0,
|
||||
"success_count": workflow.success_count or 0,
|
||||
"sys_create_datetime": workflow.sys_create_datetime,
|
||||
}
|
||||
Reference in New Issue
Block a user