Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
@@ -0,0 +1,3 @@
"""
AI 平台 API
"""
File diff suppressed because it is too large Load Diff
+295
View File
@@ -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,
}
+337
View File
@@ -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="提供商不存在")
# 构建 kwargsOllama 需要 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,
}