414 lines
14 KiB
Python
414 lines
14 KiB
Python
# =============================================================================
|
||
# 企微IT智能服务台 — AI Wingman API 路由
|
||
# =============================================================================
|
||
# 说明:坐席端 AI 智能副驾驶 API,包含以下端点:
|
||
#
|
||
# 基础能力(3 个):
|
||
# 1. POST /api/conversations/{id}/wingman/draft — 生成 AI 草稿回复
|
||
# 2. POST /api/conversations/{id}/wingman/summary — 生成会话自动摘要
|
||
# 3. POST /api/conversations/{id}/wingman/tags — 生成自动标签建议
|
||
#
|
||
# AI 辅助消息框(4 个,v_next):
|
||
# 4. POST /api/conversations/{id}/wingman/autocomplete — 自动补齐
|
||
# 5. POST /api/conversations/{id}/wingman/tone-adjust — 语气调整
|
||
# 6. POST /api/conversations/{id}/wingman/polish — 文字润色
|
||
# 7. POST /api/conversations/{id}/wingman/rewrite — 智能改写
|
||
#
|
||
# 所有端点需要坐席认证(get_current_agent)
|
||
# =============================================================================
|
||
|
||
import logging
|
||
|
||
from fastapi import APIRouter, Depends
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from app.database import get_db
|
||
from app.dependencies import dep_wingman_service
|
||
from app.models.agent import Agent
|
||
from app.models.conversation import Conversation
|
||
from app.models.message import Message
|
||
from app.services.wingman_service import WingmanService
|
||
from app.schemas.wingman_assist import (
|
||
AutocompleteRequest,
|
||
ToneAdjustRequest,
|
||
PolishRequest,
|
||
RewriteRequest,
|
||
)
|
||
from app.utils.response import ERR_NOT_FOUND, success_response
|
||
|
||
# 复用坐席认证依赖
|
||
from app.api.agents import get_current_agent
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# 创建路由器
|
||
router = APIRouter()
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# 辅助函数
|
||
# --------------------------------------------------------------------------
|
||
|
||
async def _validate_conversation(
|
||
conversation_id: str,
|
||
agent: Agent,
|
||
db: AsyncSession,
|
||
) -> Conversation:
|
||
"""验证会话存在性并返回会话对象。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
|
||
Returns:
|
||
Conversation: 会话对象
|
||
|
||
Raises:
|
||
AppException: 会话不存在
|
||
"""
|
||
stmt = select(Conversation).where(Conversation.id == conversation_id)
|
||
result = await db.execute(stmt)
|
||
conversation = result.scalars().first()
|
||
|
||
if not conversation:
|
||
raise ERR_NOT_FOUND
|
||
|
||
return conversation
|
||
|
||
|
||
async def _get_recent_messages(
|
||
conversation_id: str,
|
||
db: AsyncSession,
|
||
limit: int = 20,
|
||
) -> list[dict]:
|
||
"""获取会话最近的消息历史(转换为字典列表)。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
db: 数据库会话
|
||
limit: 获取的消息条数
|
||
|
||
Returns:
|
||
list[dict]: 消息字典列表
|
||
"""
|
||
stmt = (
|
||
select(Message)
|
||
.where(Message.conversation_id == conversation_id)
|
||
.order_by(Message.created_at.desc())
|
||
.limit(limit)
|
||
)
|
||
result = await db.execute(stmt)
|
||
messages = list(result.scalars().all())
|
||
|
||
# 按时间正序排列(最早的在前)
|
||
messages.reverse()
|
||
|
||
# 转换为字典列表
|
||
return [
|
||
{
|
||
"id": msg.id,
|
||
"sender_type": msg.sender_type,
|
||
"sender_name": msg.sender_name,
|
||
"content": msg.content,
|
||
"msg_type": msg.msg_type,
|
||
"created_at": msg.created_at.isoformat() if msg.created_at else "",
|
||
}
|
||
for msg in messages
|
||
]
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/draft
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/draft")
|
||
async def generate_draft(
|
||
conversation_id: str,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""生成 AI 草稿回复。
|
||
|
||
基于当前会话的消息历史,让 Wingman Agent 生成坐席可以采纳的草稿回复。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
agent: 当前坐席(通过认证依赖注入)
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含草稿内容、置信度和推理说明
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 从数据库读取该会话的消息历史(最近 20 条)
|
||
messages = await _get_recent_messages(conversation_id, db, limit=20)
|
||
|
||
# 3. 调用 WingmanService 生成草稿
|
||
result = await wingman_service.generate_draft(
|
||
conversation_id=conversation_id,
|
||
messages=messages,
|
||
db=db,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/summary
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/summary")
|
||
async def generate_summary(
|
||
conversation_id: str,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""生成会话自动摘要。
|
||
|
||
基于完整对话生成结构化摘要,包含问题、原因、解决方案。
|
||
通常在结单时调用。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含问题、原因、解决方案
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 从数据库读取该会话的完整消息历史(最多 50 条)
|
||
messages = await _get_recent_messages(conversation_id, db, limit=50)
|
||
|
||
# 3. 调用 WingmanService 生成摘要
|
||
result = await wingman_service.generate_summary(
|
||
conversation_id=conversation_id,
|
||
messages=messages,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/tags
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/tags")
|
||
async def suggest_tags(
|
||
conversation_id: str,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""生成自动标签建议。
|
||
|
||
基于对话内容建议标签分类,包含标签列表、分类和优先级。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含建议标签、分类和优先级
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
conversation = await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 从数据库读取该会话的消息历史(最近 20 条)
|
||
messages = await _get_recent_messages(conversation_id, db, limit=20)
|
||
|
||
# 3. 获取已有标签(用于避免重复建议)
|
||
existing_tags = {}
|
||
if hasattr(conversation, 'tags') and conversation.tags:
|
||
existing_tags = conversation.tags if isinstance(conversation.tags, dict) else {}
|
||
|
||
# 4. 调用 WingmanService 生成标签建议
|
||
result = await wingman_service.suggest_tags(
|
||
conversation_id=conversation_id,
|
||
messages=messages,
|
||
existing_tags=existing_tags,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/autocomplete
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/autocomplete")
|
||
async def autocomplete(
|
||
conversation_id: str,
|
||
request: AutocompleteRequest,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""自动补齐。
|
||
|
||
坐席输入停顿超过 800ms 后,前端请求 AI 生成下一句补齐建议。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
request: 补齐请求(当前文本、光标位置、最大长度)
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含补齐文本和置信度
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 获取最近 5 条消息作为上下文
|
||
messages = await _get_recent_messages(conversation_id, db, limit=5)
|
||
|
||
# 3. 调用 WingmanService 生成补齐
|
||
result = await wingman_service.generate_completion(
|
||
conversation_id=conversation_id,
|
||
current_text=request.current_text,
|
||
messages=messages,
|
||
max_length=request.max_length,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/tone-adjust
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/tone-adjust")
|
||
async def tone_adjust(
|
||
conversation_id: str,
|
||
request: ToneAdjustRequest,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""语气调整。
|
||
|
||
坐席选中一段文字后,选择目标语气(专业/友好/简洁),AI 对选中文字改写。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
request: 语气调整请求(选中文字、完整输入框内容、目标语气)
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含改写后的文字、语气和变更摘要
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 获取最近 5 条消息作为上下文
|
||
messages = await _get_recent_messages(conversation_id, db, limit=5)
|
||
|
||
# 3. 调用 WingmanService 进行语气调整
|
||
result = await wingman_service.adjust_tone(
|
||
conversation_id=conversation_id,
|
||
selected_text=request.selected_text,
|
||
full_text=request.full_text,
|
||
tone=request.tone,
|
||
messages=messages,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/polish
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/polish")
|
||
async def polish(
|
||
conversation_id: str,
|
||
request: PolishRequest,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""文字润色。
|
||
|
||
对坐席输入的文字进行扩写/压缩/纠错处理。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
request: 润色请求(待润色文字、操作类型、是否携带上下文)
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含润色后的文字、操作类型和变更摘要
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 可选携带对话上下文(由请求参数控制)
|
||
messages = []
|
||
if request.conversation_context:
|
||
messages = await _get_recent_messages(conversation_id, db, limit=5)
|
||
|
||
# 3. 调用 WingmanService 进行润色
|
||
result = await wingman_service.polish_text(
|
||
conversation_id=conversation_id,
|
||
text=request.text,
|
||
action=request.action,
|
||
messages=messages,
|
||
)
|
||
|
||
return success_response(data=result)
|
||
|
||
|
||
# --------------------------------------------------------------------------
|
||
# POST /api/conversations/{conversation_id}/wingman/rewrite
|
||
# --------------------------------------------------------------------------
|
||
@router.post("/conversations/{conversation_id}/wingman/rewrite")
|
||
async def rewrite(
|
||
conversation_id: str,
|
||
request: RewriteRequest,
|
||
agent: Agent = Depends(get_current_agent),
|
||
db: AsyncSession = Depends(get_db),
|
||
wingman_service: WingmanService = Depends(dep_wingman_service),
|
||
):
|
||
"""智能改写。
|
||
|
||
基于对话上下文和知识库,为坐席生成多个不同风格的备选回复。
|
||
|
||
Args:
|
||
conversation_id: 会话ID
|
||
request: 改写请求(当前输入文本、生成版本数、是否包含知识库引用)
|
||
agent: 当前坐席
|
||
db: 数据库会话
|
||
wingman_service: Wingman 服务实例
|
||
|
||
Returns:
|
||
Dict: 统一响应格式,包含多个版本的回复文本、风格标签和来源
|
||
"""
|
||
# 1. 验证坐席身份 + 会话存在性
|
||
await _validate_conversation(conversation_id, agent, db)
|
||
|
||
# 2. 获取最近 10 条消息作为上下文(改写需要更多上下文)
|
||
messages = await _get_recent_messages(conversation_id, db, limit=10)
|
||
|
||
# 3. 调用 WingmanService 进行改写
|
||
result = await wingman_service.rewrite_versions(
|
||
conversation_id=conversation_id,
|
||
current_text=request.current_text,
|
||
messages=messages,
|
||
generate_count=request.generate_count,
|
||
include_knowledge=request.include_knowledge,
|
||
)
|
||
|
||
return success_response(data=result)
|