A组认证加固: P0兜底+Token刷新+环境检测+OTP+RBRAC落地+P1日志审计+Token撤销 - 全局一致性审查通过
This commit is contained in:
@@ -1014,3 +1014,71 @@ async def ragflow_retrieval(
|
||||
return success_response(data={"error": e.message, "error_code": "config_missing"})
|
||||
except RagflowError as e:
|
||||
return success_response(data={"error": e.message, "error_code": "api_error"})
|
||||
|
||||
|
||||
# ---------- POST /api/admin/users/{employee_id}/revoke-token ----------
|
||||
@router.post("/users/{employee_id}/revoke-token")
|
||||
async def revoke_user_token(
|
||||
employee_id: str,
|
||||
admin: Agent = Depends(require_admin),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""强制撤销指定用户的登录Token(管理员操作)。
|
||||
|
||||
清除该用户在 Redis 中的所有 Token,使其被迫下线。
|
||||
同时记录审计日志。
|
||||
|
||||
Args:
|
||||
employee_id: 要撤销 Token 的用户 ID(企微 UserID)
|
||||
|
||||
Returns:
|
||||
撤销结果
|
||||
"""
|
||||
from app.dependencies import get_redis
|
||||
from app.services.audit_log_service import record_audit_log
|
||||
|
||||
redis_client = await get_redis()
|
||||
|
||||
# 搜索可能的 Token key 模式
|
||||
# 1. user:token:* - 统一格式
|
||||
# 2. agent:token:* - 坐席端
|
||||
# 3. employee:token:* - 员工端
|
||||
revoked_count = 0
|
||||
patterns = ["user:token:*", "agent:token:*", "employee:token:*"]
|
||||
|
||||
for pattern in patterns:
|
||||
cursor = 0
|
||||
while True:
|
||||
cursor, keys = await redis_client.scan(cursor, match=pattern, count=100)
|
||||
for key in keys:
|
||||
token_data = await redis_client.get(key)
|
||||
if token_data:
|
||||
try:
|
||||
import json
|
||||
data = json.loads(token_data)
|
||||
if data.get("employee_id") == employee_id:
|
||||
await redis_client.delete(key)
|
||||
revoked_count += 1
|
||||
logger.info(f"撤销 Token: key={key}, employee_id={employee_id}")
|
||||
except (json.JSONDecodeError, Exception):
|
||||
pass
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
# 记录审计日志
|
||||
await record_audit_log(
|
||||
db=db,
|
||||
employee_id=admin.employee_id,
|
||||
action="revoke_token",
|
||||
resource="user",
|
||||
resource_id=employee_id,
|
||||
details={"revoked_count": revoked_count, "operator": admin.employee_id},
|
||||
result="success",
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
return success_response(data={
|
||||
"employee_id": employee_id,
|
||||
"revoked_count": revoked_count,
|
||||
"message": f"已撤销 {revoked_count} 个 Token" if revoked_count > 0 else "未找到该用户的有效 Token",
|
||||
})
|
||||
|
||||
@@ -39,9 +39,11 @@ import logging
|
||||
from typing import Optional
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
from fastapi import APIRouter, Depends, Path
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.dependencies import dep_redis, get_current_user, UserInfo
|
||||
from app.schemas.qrcode import (
|
||||
QrcodeConfirmRequest,
|
||||
@@ -139,30 +141,85 @@ async def poll_qrcode(
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# POST /api/auth_qrcode/scan — 企微 OAuth code 回调
|
||||
# GET|POST /api/auth_qrcode/scan — 企微 OAuth code 回调
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/scan", response_model=None)
|
||||
@router.api_route("/scan", methods=["GET", "POST"], response_model=None)
|
||||
async def scan_qrcode(
|
||||
body: QrcodeScanRequest,
|
||||
body: Optional[QrcodeScanRequest] = None,
|
||||
ticket: Optional[str] = Query(None, description="扫码登录票据(兼容旧参数名)"),
|
||||
state: Optional[str] = Query(None, description="扫码登录票据(企微 OAuth state 标准参数名)"),
|
||||
code: Optional[str] = Query(None, description="企微 OAuth 授权码"),
|
||||
redis_client: aioredis.Redis = Depends(dep_redis),
|
||||
):
|
||||
"""处理企微 OAuth2 扫码回调。
|
||||
|
||||
企微 OAuth2 标准回调走 **GET** 带 query 参数 `?code=xxx&state=<ticket>`,
|
||||
本端点同时支持 GET 和 POST(POST 兼容内部调用 / 旧前端代码)。
|
||||
|
||||
GET 模式 (企微 OAuth2 标准回调):
|
||||
- ticket ← query.state
|
||||
- code ← query.code
|
||||
- 自动 302 跳转到 /itdesk/ 或 /itadmin/ 或 /itagent/(按角色)
|
||||
POST 模式 (内部调用):
|
||||
- ticket ← body.ticket
|
||||
- code ← body.code
|
||||
|
||||
无需鉴权(此端点被企微服务器回调,带 code + ticket)。
|
||||
用 code 换取企微 userid,然后写 Redis scan:{ticket} 等待 confirm 端点。
|
||||
|
||||
dev 模式: code 形如 "dev:dev-user-001",跳过企微 API 调用。
|
||||
|
||||
Args:
|
||||
body: 包含 ticket 和 code
|
||||
|
||||
Returns:
|
||||
Dict: 统一响应格式,data 字段是 QrcodeScanResponse
|
||||
"""
|
||||
try:
|
||||
service = _get_qrcode_service(redis_client)
|
||||
result = await service.process_scan(ticket=body.ticket, code=body.code)
|
||||
# 1. 解析参数:POST 用 body,GET 用 query
|
||||
if body is not None:
|
||||
final_ticket = body.ticket
|
||||
final_code = body.code
|
||||
else:
|
||||
# 优先用 state(企微 OAuth 标准),回退到 ticket(兼容旧调用)
|
||||
final_ticket = state or ticket
|
||||
final_code = code
|
||||
|
||||
if not final_ticket or not final_code:
|
||||
logger.warning(f"扫码参数缺失: ticket={final_ticket!r}, code={final_code!r}")
|
||||
raise AppException(1000, "缺少 ticket 或 code 参数")
|
||||
|
||||
service = _get_qrcode_service(redis_client)
|
||||
result = await service.process_scan(ticket=final_ticket, code=final_code)
|
||||
|
||||
# GET 请求(企微 OAuth 回调)→ 重定向到前端选择页
|
||||
# 因为企微 OAuth 流程不在这个端点完成最终登录,只标记 scanned,
|
||||
# 等用户在坐席端点 confirm 后才能拿到 token。
|
||||
# 但企微 WebView 期望看到跳转后的页面,所以这里给个提示页。
|
||||
from fastapi.responses import HTMLResponse
|
||||
if code is not None and ticket is not None and body is None:
|
||||
# GET 模式:渲染一个 "扫码成功" 的 HTML 提示页 + 引导用户到登录页
|
||||
html = f"""<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>扫码成功</title>
|
||||
<style>
|
||||
body {{ font-family: -apple-system, BlinkMacSystemFont, sans-serif; background: #0f172a; color: #e2e8f0; display: flex; align-items: center; justify-content: center; min-height: 100vh; margin: 0; padding: 20px; }}
|
||||
.card {{ background: #1e293b; border-radius: 16px; padding: 40px 32px; max-width: 400px; text-align: center; box-shadow: 0 10px 30px rgba(0,0,0,0.3); }}
|
||||
h1 {{ color: #34d399; margin: 0 0 16px 0; font-size: 24px; }}
|
||||
p {{ color: #94a3b8; margin: 8px 0; line-height: 1.6; }}
|
||||
.ico {{ font-size: 56px; margin-bottom: 16px; }}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="card">
|
||||
<div class="ico">✅</div>
|
||||
<h1>扫码成功</h1>
|
||||
<p>请在刚才打开登录页的浏览器中</p>
|
||||
<p>点击 <strong style="color:#60a5fa">「确认登录」</strong> 按钮完成登录</p>
|
||||
<p style="margin-top:24px;font-size:13px;color:#64748b">本页可关闭</p>
|
||||
</div>
|
||||
</body>
|
||||
</html>"""
|
||||
return HTMLResponse(content=html, status_code=200)
|
||||
|
||||
# POST 模式:返回 JSON
|
||||
return success_response(data={
|
||||
"success": result["success"],
|
||||
"message": result["message"],
|
||||
@@ -173,7 +230,7 @@ async def scan_qrcode(
|
||||
logger.warning(f"扫码业务错误: {ve}")
|
||||
raise AppException(1003, str(ve))
|
||||
except Exception as e:
|
||||
logger.error(f"扫码处理异常: ticket={body.ticket[:8]}..., error={e}", exc_info=True)
|
||||
logger.error(f"扫码处理异常: error={e}", exc_info=True)
|
||||
raise AppException(1005, f"扫码处理失败: {str(e)}")
|
||||
|
||||
|
||||
@@ -185,6 +242,7 @@ async def confirm_qrcode(
|
||||
body: QrcodeConfirmRequest,
|
||||
current_user: UserInfo = Depends(get_current_user),
|
||||
redis_client: aioredis.Redis = Depends(dep_redis),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""处理当前已登录坐席的扫码确认授权。
|
||||
|
||||
@@ -213,6 +271,24 @@ async def confirm_qrcode(
|
||||
otp_code=body.otp_code,
|
||||
)
|
||||
|
||||
# 记录扫码登录日志(成功)
|
||||
from app.services.audit_log_service import record_audit_log
|
||||
await record_audit_log(
|
||||
db=db,
|
||||
employee_id=result["employee_id"],
|
||||
action="qrcode_login",
|
||||
resource="auth",
|
||||
resource_id=result["employee_id"],
|
||||
details={
|
||||
"name": result["name"],
|
||||
"roles": result["roles"],
|
||||
"confirmed_by": current_user.employee_id,
|
||||
"login_method": "qrcode_confirm",
|
||||
},
|
||||
result="success",
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
return success_response(data={
|
||||
"token": result["token"],
|
||||
"employee_id": result["employee_id"],
|
||||
|
||||
@@ -20,6 +20,7 @@
|
||||
# =============================================================================
|
||||
|
||||
import logging
|
||||
import re
|
||||
import secrets
|
||||
import urllib.parse
|
||||
from datetime import datetime, timedelta
|
||||
@@ -35,6 +36,7 @@ from app.database import get_db
|
||||
from app.models.role import Role
|
||||
from app.models.user_role import UserRole
|
||||
from app.services.wecom_service import WecomService
|
||||
from app.services.audit_log_service import record_audit_log
|
||||
from app.utils.response import AppException
|
||||
from app.dependencies import get_redis
|
||||
|
||||
@@ -46,6 +48,15 @@ router = APIRouter(prefix="/auth_wecom", tags=["企微 SSO"])
|
||||
OAUTH_STATE_TTL = 300
|
||||
# SSO token 长度
|
||||
SSO_TOKEN_BYTES = 32
|
||||
# Token TTL 常量(8小时)
|
||||
TOKEN_TTL_SECONDS = 8 * 60 * 60 # 8小时
|
||||
# 企微 API 超时设置(秒)
|
||||
WECOM_API_TIMEOUT = 10
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# 企微环境检测(使用统一工具模块)
|
||||
# --------------------------------------------------------------------------
|
||||
from app.utils.wecom_auth import require_wecom_ua as _require_wework_ua
|
||||
|
||||
|
||||
def _sso_enabled() -> bool:
|
||||
@@ -101,6 +112,9 @@ async def sso_init(
|
||||
Args:
|
||||
next: 登录成功后跳转路径,如 /itdesk/ /itagent/ /itadmin/
|
||||
"""
|
||||
# 后端第二道防线:非企微环境拒绝授权
|
||||
_require_wework_ua(request)
|
||||
|
||||
if not _sso_enabled():
|
||||
raise AppException(1001, "企微 SSO 未启用, 请用扫码登录")
|
||||
|
||||
@@ -116,103 +130,229 @@ async def sso_init(
|
||||
str(state_payload).encode("utf-8"),
|
||||
)
|
||||
|
||||
# 2. 拼企微 OAuth URL
|
||||
# 2. 拼企微 OAuth URL(回调URL中包含next参数,用于state失效时仍能知道目标路径)
|
||||
callback_url = _get_oauth_callback_url(request)
|
||||
oauth_url = _build_oauth_url(state, callback_url)
|
||||
# 在回调URL中添加next参数
|
||||
separator = "&" if "?" in callback_url else "?"
|
||||
callback_url_with_next = f"{callback_url}{separator}next={urllib.parse.quote(next)}"
|
||||
oauth_url = _build_oauth_url(state, callback_url_with_next)
|
||||
|
||||
logger.info(f"SSO init: state={state[:8]}..., next={next}")
|
||||
return RedirectResponse(url=oauth_url, status_code=302)
|
||||
|
||||
|
||||
def _get_error_redirect_url(error_code: str, error_msg: str, next_path: str = "/itdesk/") -> str:
|
||||
"""生成 OAuth 错误重定向 URL。
|
||||
|
||||
异常时重定向到前端错误页面(ErrorPage),而不是返回 JSON 错误。
|
||||
ErrorPage 读取 code 和 message 参数显示友好错误提示。
|
||||
|
||||
Args:
|
||||
error_code: 错误码
|
||||
error_msg: 错误信息
|
||||
next_path: 原始请求的目标路径(保留但不再用于决定重定向)
|
||||
"""
|
||||
import os
|
||||
base = getattr(settings, "wecom_sso_callback_base", None)
|
||||
if not base:
|
||||
base = os.getenv("WECOM_SSO_CALLBACK_BASE", "https://itsupport.servyou.com.cn")
|
||||
|
||||
# 重定向到 Portal 的 ErrorPage,带错误参数
|
||||
# ErrorPage 期望格式:?code=xxx&message=yyy
|
||||
return f"{base.rstrip('/')}/itportal/error?code={error_code}&message={urllib.parse.quote(error_msg)}"
|
||||
|
||||
|
||||
@router.get("/sso/callback")
|
||||
async def sso_callback(
|
||||
code: str = Query(..., description="企微 OAuth2 授权 code"),
|
||||
state: str = Query(..., description="防 CSRF state"),
|
||||
request: Request,
|
||||
code: Optional[str] = Query(None, description="企微 OAuth2 授权 code"),
|
||||
state: Optional[str] = Query(None, description="防 CSRF state"),
|
||||
errcode: Optional[int] = Query(None, description="企微 OAuth 错误码"),
|
||||
errmsg: Optional[str] = Query(None, description="企微 OAuth 错误信息"),
|
||||
next: Optional[str] = Query(None, description="原始请求的目标路径(可选,用于错误时重定向)"),
|
||||
redis_client = Depends(get_redis),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""企微 OAuth 回调: 用 code 换 userid → 查 role → 生成 token → 跳 next。"""
|
||||
# 1. 校验 state(防 CSRF)
|
||||
state_key = f"wecom_sso:state:{state}"
|
||||
state_raw = await redis_client.get(state_key)
|
||||
if not state_raw:
|
||||
raise AppException(1002, "SSO state 已过期或无效, 请重新进入")
|
||||
"""企微 OAuth 回调: 用 code 换 userid → 查 role → 生成 token → 跳 next。
|
||||
|
||||
# 删除 state(一次性)
|
||||
await redis_client.delete(state_key)
|
||||
异常时重定向到前端错误页面,避免白屏。
|
||||
所有未处理的异常都会记录详细日志(包含 traceback)。
|
||||
|
||||
import ast
|
||||
state_data = ast.literal_eval(state_raw.decode("utf-8"))
|
||||
next_path = state_data.get("next", "/itdesk/")
|
||||
Args:
|
||||
next: 原始请求的目标路径,用于错误重定向。如果 state 验证失败,使用此参数决定重定向位置。
|
||||
"""
|
||||
import traceback
|
||||
|
||||
# 默认 next 路径
|
||||
next_path = next or "/itdesk/"
|
||||
|
||||
# 2. 用 code 换 userid
|
||||
wecom = WecomService(redis_client)
|
||||
try:
|
||||
oauth_info = await wecom.get_oauth_user_info(code)
|
||||
user_id = oauth_info.get("userid", "")
|
||||
if not user_id:
|
||||
raise AppException(1003, "企微 OAuth 返回 userid 为空")
|
||||
# 后端第二道防线:非企微环境拒绝回调
|
||||
_require_wework_ua(request)
|
||||
|
||||
user_info = await wecom.get_user_info(user_id)
|
||||
name = user_info.get("name", user_id)
|
||||
except Exception as e:
|
||||
logger.error(f"SSO callback 调企微 API 失败: code={code[:8]}..., error={e}")
|
||||
raise AppException(1004, f"企微身份识别失败: {str(e)}")
|
||||
finally:
|
||||
# 0. 处理企微返回的错误(用户拒绝授权等)
|
||||
if errcode is not None:
|
||||
msg = errmsg or "用户取消授权或授权失败"
|
||||
logger.warning(f"SSO callback 企微返回错误: errcode={errcode}, errmsg={errmsg}")
|
||||
return RedirectResponse(url=_get_error_redirect_url(f"wecom_{errcode}", msg, next_path), status_code=302)
|
||||
|
||||
# 1. 校验必要参数
|
||||
if not code or not state:
|
||||
logger.warning(f"SSO callback 缺少必要参数: code={bool(code)}, state={bool(state)}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("missing_params", "授权参数不完整,请重试", next_path), status_code=302)
|
||||
|
||||
# 2. 校验 state(防 CSRF)
|
||||
state_key = f"wecom_sso:state:{state}"
|
||||
try:
|
||||
await wecom.close()
|
||||
except Exception:
|
||||
pass
|
||||
state_raw = await redis_client.get(state_key)
|
||||
except Exception as e:
|
||||
logger.error(f"SSO callback Redis 获取 state 失败: {e}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("redis_error", "服务暂不可用,请稍后重试", next_path), status_code=302)
|
||||
|
||||
# 3. 查 role (user/agent/admin)
|
||||
role_stmt = (
|
||||
select(Role)
|
||||
.join(UserRole, Role.id == UserRole.role_id)
|
||||
.where(UserRole.employee_id == user_id)
|
||||
)
|
||||
role_result = await db.execute(role_stmt)
|
||||
roles = role_result.scalars().all()
|
||||
if not state_raw:
|
||||
logger.warning(f"SSO callback state 过期: state={state[:8]}...")
|
||||
return RedirectResponse(url=_get_error_redirect_url("state_expired", "授权已过期,请重新进入", next_path), status_code=302)
|
||||
|
||||
if not roles:
|
||||
# 没有绑定角色: 跳"无权限"页
|
||||
logger.warning(f"SSO: user_id={user_id} 没绑定任何角色")
|
||||
return RedirectResponse(url=f"/itdesk/no-role?user_id={user_id}", status_code=302)
|
||||
# 删除 state(一次性)
|
||||
try:
|
||||
await redis_client.delete(state_key)
|
||||
except Exception as e:
|
||||
logger.warning(f"SSO callback 删除 state 失败: {e}") # 不阻塞流程
|
||||
|
||||
# 4. 选最高权限角色 (admin > agent > user)
|
||||
role_priority = {"admin": 3, "agent": 2, "user": 1}
|
||||
best_role = max(roles, key=lambda r: role_priority.get(r.name, 0))
|
||||
role_name = best_role.name
|
||||
# 解析 state 数据(添加异常处理)
|
||||
import ast
|
||||
import json
|
||||
try:
|
||||
state_data = json.loads(state_raw.decode("utf-8"))
|
||||
except (json.JSONDecodeError, AttributeError) as e:
|
||||
# 兼容旧格式(使用 ast.literal_eval)
|
||||
try:
|
||||
state_data = ast.literal_eval(state_raw.decode("utf-8"))
|
||||
except (ValueError, SyntaxError) as e2:
|
||||
logger.error(f"SSO callback state 解析失败: {e2}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("state_invalid", "授权信息无效,请重新进入", next_path), status_code=302)
|
||||
|
||||
# 5. 生成 SSO token(随机 + Redis 存 8 小时)
|
||||
sso_token = secrets.token_urlsafe(SSO_TOKEN_BYTES)
|
||||
sso_payload = {
|
||||
"user_id": user_id,
|
||||
"name": name,
|
||||
"role": role_name,
|
||||
"created_at": datetime.now().isoformat(),
|
||||
}
|
||||
import json
|
||||
await redis_client.setex(
|
||||
f"wecom_sso:token:{sso_token}",
|
||||
8 * 3600, # 8 小时
|
||||
json.dumps(sso_payload, ensure_ascii=False).encode("utf-8"),
|
||||
)
|
||||
# 优先使用 state 中存储的 next,fallback 到 URL 参数
|
||||
next_path = state_data.get("next", next_path)
|
||||
|
||||
# 6. 跳转到 next + token
|
||||
separator = "&" if "?" in next_path else "?"
|
||||
redirect_url = f"{next_path}{separator}sso_token={sso_token}"
|
||||
# 3. 用 code 换 userid
|
||||
wecom = WecomService(redis_client)
|
||||
try:
|
||||
oauth_info = await wecom.get_oauth_user_info(code)
|
||||
user_id = oauth_info.get("userid", "")
|
||||
if not user_id:
|
||||
logger.warning("SSO callback 企微返回 userid 为空")
|
||||
return RedirectResponse(url=_get_error_redirect_url("empty_userid", "无法获取您的企业微信身份,请重试", next_path), status_code=302)
|
||||
|
||||
logger.info(f"SSO 成功: user_id={user_id}, role={role_name}, next={next_path}")
|
||||
return RedirectResponse(url=redirect_url, status_code=302)
|
||||
user_info = await wecom.get_user_info(user_id)
|
||||
name = user_info.get("name", user_id)
|
||||
except Exception as e:
|
||||
logger.error(f"SSO callback 调企微 API 失败: code={code[:8]}..., error={e}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("api_failed", f"企业微信服务异常: {str(e)}"), status_code=302)
|
||||
finally:
|
||||
try:
|
||||
await wecom.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 3. 查 role (user/agent/admin)
|
||||
try:
|
||||
role_stmt = (
|
||||
select(Role)
|
||||
.join(UserRole, Role.id == UserRole.role_id)
|
||||
.where(UserRole.employee_id == user_id)
|
||||
)
|
||||
role_result = await db.execute(role_stmt)
|
||||
roles = role_result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error(f"SSO callback 查询角色失败: {e}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("db_error", "服务暂不可用,请稍后重试", next_path), status_code=302)
|
||||
|
||||
if not roles:
|
||||
# 没有绑定角色: 跳"无权限"页
|
||||
logger.warning(f"SSO: user_id={user_id} 没绑定任何角色")
|
||||
return RedirectResponse(url=f"/itdesk/no-role?user_id={user_id}", status_code=302)
|
||||
|
||||
# 4. 选最高权限角色 (admin > agent > user)
|
||||
role_priority = {"admin": 3, "agent": 2, "user": 1}
|
||||
best_role = max(roles, key=lambda r: role_priority.get(r.name, 0))
|
||||
role_name = best_role.name
|
||||
|
||||
# 5. 生成 SSO token(随机 + Redis 存 8 小时)
|
||||
sso_token = secrets.token_urlsafe(SSO_TOKEN_BYTES)
|
||||
sso_payload = {
|
||||
"user_id": user_id,
|
||||
"name": name,
|
||||
"role": role_name,
|
||||
"created_at": datetime.now().isoformat(),
|
||||
}
|
||||
import json
|
||||
try:
|
||||
await redis_client.setex(
|
||||
f"wecom_sso:token:{sso_token}",
|
||||
TOKEN_TTL_SECONDS,
|
||||
json.dumps(sso_payload, ensure_ascii=False).encode("utf-8"),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"SSO callback 存储 token 失败: {e}")
|
||||
return RedirectResponse(url=_get_error_redirect_url("redis_error", "服务暂不可用,请稍后重试", next_path), status_code=302)
|
||||
|
||||
# 6. 记录登录日志
|
||||
try:
|
||||
await record_audit_log(
|
||||
db=db,
|
||||
employee_id=user_id,
|
||||
action="sso_login",
|
||||
resource="auth",
|
||||
resource_id=user_id,
|
||||
details={"name": name, "role": role_name, "login_method": "wecom_sso"},
|
||||
result="success",
|
||||
request=None, # callback 请求没有直接可用的 request 对象
|
||||
)
|
||||
await db.commit()
|
||||
except Exception as e:
|
||||
logger.warning(f"SSO callback 记录登录日志失败: {e}") # 不阻塞登录流程
|
||||
|
||||
# 7. 跳转到 next + token
|
||||
separator = "&" if "?" in next_path else "?"
|
||||
redirect_url = f"{next_path}{separator}sso_token={sso_token}"
|
||||
|
||||
logger.info(f"SSO 成功: user_id={user_id}, role={role_name}, next={next_path}")
|
||||
return RedirectResponse(url=redirect_url, status_code=302)
|
||||
|
||||
except Exception as e:
|
||||
# 捕获所有未处理的异常,记录详细日志(包含 traceback)并重定向到错误页
|
||||
error_details = {
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"code": code[:8] + "..." if code else None,
|
||||
"state": state[:8] + "..." if state else None,
|
||||
"next": next_path,
|
||||
}
|
||||
logger.error(
|
||||
f"SSO callback 未处理的异常: {error_details}\n"
|
||||
f"traceback: {traceback.format_exc()}"
|
||||
)
|
||||
return RedirectResponse(
|
||||
url=_get_error_redirect_url("oauth_failed", "登录过程出现异常,请重试"),
|
||||
status_code=302
|
||||
)
|
||||
|
||||
|
||||
@router.get("/sso/verify")
|
||||
async def sso_verify(
|
||||
request: Request,
|
||||
sso_token: str = Query(..., description="SSO token"),
|
||||
redis_client = Depends(get_redis),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""前端用 SSO token 换用户身份(token 一次性使用,用完删除)。"""
|
||||
"""前端用 SSO token 换用户身份(token 一次性使用,用完删除)。
|
||||
|
||||
后端第二道防线:非企微环境拒绝验证。
|
||||
"""
|
||||
# 后端第二道防线:非企微环境拒绝验证
|
||||
_require_wework_ua(request)
|
||||
|
||||
import json
|
||||
token_raw = await redis_client.get(f"wecom_sso:token:{sso_token}")
|
||||
if not token_raw:
|
||||
@@ -226,3 +366,162 @@ async def sso_verify(
|
||||
"code": 0,
|
||||
"data": payload,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/refresh")
|
||||
async def refresh_token(
|
||||
token: str = Query(..., description="当前 Bearer token"),
|
||||
redis_client = Depends(get_redis),
|
||||
):
|
||||
"""刷新 Token TTL。
|
||||
|
||||
前端在 Token 过期前 5 分钟自动调用此接口,实现静默刷新。
|
||||
如果 Token 无效或已过期,返回 401 错误。
|
||||
|
||||
Returns:
|
||||
刷新成功:{ code: 0, data: { token: "新token", expires_in: 28800 } }
|
||||
"""
|
||||
import json
|
||||
|
||||
# 1. 尝试统一格式 Token
|
||||
token_key = f"user:token:{token}"
|
||||
token_data_raw = await redis_client.get(token_key)
|
||||
|
||||
if token_data_raw:
|
||||
try:
|
||||
user_info = json.loads(token_data_raw)
|
||||
# 更新最后活跃时间
|
||||
user_info["last_active"] = datetime.now().isoformat()
|
||||
|
||||
# 延长 TTL(重新设置 8 小时)
|
||||
await redis_client.setex(
|
||||
token_key,
|
||||
TOKEN_TTL_SECONDS,
|
||||
json.dumps(user_info, ensure_ascii=False),
|
||||
)
|
||||
|
||||
logger.info(f"Token 刷新成功: employee_id={user_info.get('employee_id')}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token, # 复用同一个 token,只延长 TTL
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 2. 尝试旧格式 Token (employee:token)
|
||||
employee_key = f"employee:token:{token}"
|
||||
employee_id = await redis_client.get(employee_key)
|
||||
if employee_id:
|
||||
# 延长 TTL
|
||||
await redis_client.expire(employee_key, TOKEN_TTL_SECONDS)
|
||||
logger.info(f"Token 刷新成功(employee): employee_id={employee_id}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token,
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
|
||||
# 3. 尝试旧格式 Token (agent:token)
|
||||
agent_key = f"agent:token:{token}"
|
||||
agent_id = await redis_client.get(agent_key)
|
||||
if agent_id:
|
||||
await redis_client.expire(agent_key, TOKEN_TTL_SECONDS)
|
||||
logger.info(f"Token 刷新成功(agent): agent_id={agent_id}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token,
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
|
||||
# Token 无效或已过期
|
||||
logger.warning(f"Token 刷新失败: token 不存在或已过期")
|
||||
raise AppException(401, "Token 已过期,请重新登录")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# 别名路由:支持前端 /api/auth/refresh 调用(与 /api/auth_wecom/refresh 等效)
|
||||
# --------------------------------------------------------------------------
|
||||
# 前端 H5/坐席/管理后台调用 /api/auth/refresh,后端响应 /api/auth_wecom/refresh
|
||||
# 为兼容前端习惯,添加此别名路由
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
# 创建别名路由器(无 prefix)
|
||||
alias_router = APIRouter(tags=["认证"])
|
||||
|
||||
|
||||
@alias_router.post("/auth/refresh")
|
||||
async def refresh_token_alias(
|
||||
token: str = Query(..., description="当前 Bearer token"),
|
||||
redis_client = Depends(get_redis),
|
||||
):
|
||||
"""Token 刷新接口别名。
|
||||
|
||||
前端调用 /api/auth/refresh,后端实际处理逻辑与 /api/auth_wecom/refresh 相同。
|
||||
这是为了兼容前端的调用习惯。
|
||||
|
||||
Returns:
|
||||
刷新成功:{ code: 0, data: { token: "新token", expires_in: 28800 } }
|
||||
"""
|
||||
import json
|
||||
|
||||
# 1. 尝试统一格式 Token
|
||||
token_key = f"user:token:{token}"
|
||||
token_data_raw = await redis_client.get(token_key)
|
||||
|
||||
if token_data_raw:
|
||||
try:
|
||||
user_info = json.loads(token_data_raw)
|
||||
user_info["last_active"] = datetime.now().isoformat()
|
||||
await redis_client.setex(
|
||||
token_key,
|
||||
TOKEN_TTL_SECONDS,
|
||||
json.dumps(user_info, ensure_ascii=False),
|
||||
)
|
||||
logger.info(f"Token 刷新成功(alias): employee_id={user_info.get('employee_id')}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token,
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 2. 尝试旧格式 Token
|
||||
employee_key = f"employee:token:{token}"
|
||||
employee_id = await redis_client.get(employee_key)
|
||||
if employee_id:
|
||||
await redis_client.expire(employee_key, TOKEN_TTL_SECONDS)
|
||||
logger.info(f"Token 刷新成功(alias employee): employee_id={employee_id}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token,
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
|
||||
# 3. 尝试 agent token
|
||||
agent_key = f"agent:token:{token}"
|
||||
agent_id = await redis_client.get(agent_key)
|
||||
if agent_id:
|
||||
await redis_client.expire(agent_key, TOKEN_TTL_SECONDS)
|
||||
logger.info(f"Token 刷新成功(alias agent): agent_id={agent_id}")
|
||||
return {
|
||||
"code": 0,
|
||||
"data": {
|
||||
"token": token,
|
||||
"expires_in": TOKEN_TTL_SECONDS,
|
||||
},
|
||||
}
|
||||
|
||||
logger.warning(f"Token 刷新失败(alias): token 不存在或已过期")
|
||||
raise AppException(401, "Token 已过期,请重新登录")
|
||||
|
||||
@@ -30,6 +30,7 @@ from app.schemas.conversation import (
|
||||
ConversationStatusUpdate,
|
||||
InviteParticipantRequest,
|
||||
JoinConversationRequest,
|
||||
UpdateTagsRequest,
|
||||
)
|
||||
from app.services.session_service import SessionService
|
||||
from app.services.wecom_service import WecomService
|
||||
@@ -38,6 +39,9 @@ from app.utils.response import AppException, success_response
|
||||
# 坐席认证依赖(从 agents.py 导入)
|
||||
from app.api.agents import get_current_agent
|
||||
|
||||
# RBAC 权限装饰器
|
||||
from app.dependencies import require_role, require_permission
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 创建路由器
|
||||
@@ -48,6 +52,7 @@ router = APIRouter()
|
||||
# GET /api/conversations — 获取坐席会话列表(全局可见)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.get("/conversations")
|
||||
@require_permission("conversation", "read", "all")
|
||||
async def list_conversations(
|
||||
status: Optional[str] = Query(None, description="按状态过滤: ai_handling/queued/serving/resolved"),
|
||||
agent_id: Optional[str] = Query(None, description="按坐席ID过滤"),
|
||||
@@ -142,6 +147,7 @@ async def list_conversations(
|
||||
# GET /api/conversations/{id} — 获取会话详情
|
||||
# --------------------------------------------------------------------------
|
||||
@router.get("/conversations/{conversation_id}")
|
||||
@require_permission("conversation", "read", "all")
|
||||
async def get_conversation(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -166,6 +172,7 @@ async def get_conversation(
|
||||
# POST /api/conversations/{id}/assign — 坐席接单
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/assign")
|
||||
@require_permission("conversation", "update", "all")
|
||||
async def assign_conversation(
|
||||
conversation_id: str,
|
||||
body: ConversationAssign,
|
||||
@@ -216,6 +223,7 @@ async def assign_conversation(
|
||||
# POST /api/conversations/{id}/resolve — 结单
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/resolve")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def resolve_conversation(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -259,6 +267,7 @@ async def resolve_conversation(
|
||||
# POST /api/conversations/{id}/pin — 置顶/取消置顶
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/pin")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def toggle_pin(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -285,6 +294,7 @@ async def toggle_pin(
|
||||
# POST /api/conversations/{id}/todo — 代办/取消代办
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/todo")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def toggle_todo(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -311,6 +321,7 @@ async def toggle_todo(
|
||||
# POST /api/conversations/{id}/transfer — 转接
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/transfer")
|
||||
@require_permission("conversation", "update", "all")
|
||||
async def transfer_conversation(
|
||||
conversation_id: str,
|
||||
body: ConversationAssign,
|
||||
@@ -342,6 +353,7 @@ async def transfer_conversation(
|
||||
# POST /api/conversations/{id}/grab — 接手会话(抢单)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/grab")
|
||||
@require_permission("conversation", "update", "all")
|
||||
async def grab_conversation(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -439,6 +451,7 @@ async def grab_conversation(
|
||||
# POST /api/conversations/{id}/invite — 摇人(邀请坐席协作)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/invite")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def invite_collaborator(
|
||||
conversation_id: str,
|
||||
body: ConversationInvite,
|
||||
@@ -484,6 +497,7 @@ async def invite_collaborator(
|
||||
# POST /api/conversations/{id}/leave — 退出协作
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/leave")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def leave_collaboration(
|
||||
conversation_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
@@ -529,6 +543,7 @@ async def leave_collaboration(
|
||||
# POST /api/conversations/{id}/invite-participant — 邀请员工/部门加入会话
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/invite-participant")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def invite_participant(
|
||||
conversation_id: str,
|
||||
body: InviteParticipantRequest,
|
||||
@@ -587,6 +602,7 @@ async def invite_participant(
|
||||
# POST /api/conversations/{id}/join — 被邀请人加入会话
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/join")
|
||||
@require_permission("conversation", "update", "all")
|
||||
async def join_conversation(
|
||||
conversation_id: str,
|
||||
body: JoinConversationRequest,
|
||||
@@ -622,6 +638,7 @@ async def join_conversation(
|
||||
# DELETE /api/conversations/{id}/participants/{user_id} — 移除参与者
|
||||
# --------------------------------------------------------------------------
|
||||
@router.delete("/conversations/{conversation_id}/participants/{user_id}")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def remove_participant(
|
||||
conversation_id: str,
|
||||
user_id: str,
|
||||
@@ -659,6 +676,7 @@ async def remove_participant(
|
||||
# POST /api/conversations/{id}/leave-participant — 参与者主动退出
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/leave-participant")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def leave_as_participant(
|
||||
conversation_id: str,
|
||||
body: JoinConversationRequest,
|
||||
@@ -686,3 +704,60 @@ async def leave_as_participant(
|
||||
|
||||
response_data = ConversationResponse.model_validate(conversation).model_dump()
|
||||
return success_response(data=response_data)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# POST /api/conversations/{conversation_id}/tags — 保存会话标签
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/tags")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def update_conversation_tags(
|
||||
conversation_id: str,
|
||||
body: UpdateTagsRequest,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_agent: Agent = Depends(get_current_agent),
|
||||
):
|
||||
"""保存会话标签。
|
||||
|
||||
坐席可以为会话添加/更新标签,如问题分类、优先级、情绪状态等。
|
||||
标签以 JSON 形式存储在会话的 tags 字段中。
|
||||
|
||||
Args:
|
||||
conversation_id: 会话ID
|
||||
body: 标签更新请求,包含 tags 字典
|
||||
current_agent: 当前坐席(通过认证依赖注入)
|
||||
db: 数据库会话
|
||||
|
||||
Returns:
|
||||
更新后的会话详情
|
||||
"""
|
||||
# 1. 验证会话存在性
|
||||
stmt = select(Conversation).where(Conversation.id == conversation_id)
|
||||
result = await db.execute(stmt)
|
||||
conversation = result.scalars().first()
|
||||
|
||||
if not conversation:
|
||||
raise AppException("会话不存在", code=404)
|
||||
|
||||
# 2. 合并现有标签(如果有)
|
||||
existing_tags = {}
|
||||
if conversation.tags:
|
||||
existing_tags = (
|
||||
dict(conversation.tags) if isinstance(conversation.tags, dict) else {}
|
||||
)
|
||||
|
||||
# 3. 合并新旧标签(body.tags 覆盖同名 key)
|
||||
merged_tags = {**existing_tags, **body.tags}
|
||||
|
||||
# 4. 保存到数据库
|
||||
conversation.tags = merged_tags
|
||||
await db.commit()
|
||||
await db.refresh(conversation)
|
||||
|
||||
logger.info(
|
||||
f"坐席 {current_agent.id} 更新会话 {conversation_id} 标签: {merged_tags}"
|
||||
)
|
||||
|
||||
# 5. 返回更新后的会话
|
||||
response_data = ConversationResponse.model_validate(conversation).model_dump()
|
||||
return success_response(data=response_data)
|
||||
|
||||
+152
-2
@@ -10,17 +10,19 @@
|
||||
# 6. POST /api/conversations/{id}/mark-read — 标记已读
|
||||
# 7. POST /api/messages/image — 上传图片
|
||||
# 8. POST /api/messages/file — 上传文件
|
||||
# 9. GET /api/conversations/{id}/messages/search — 搜索消息(MSG-P1-04)
|
||||
# 消息发送需同时:存数据库 + 调用企微API发送给员工
|
||||
# =============================================================================
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import APIRouter, Depends, File, Query, UploadFile
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy import select, update, or_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
@@ -29,7 +31,12 @@ from app.models.conversation import Conversation
|
||||
from app.models.message import Message
|
||||
from app.schemas.message import MessageCreate, MessageResponse
|
||||
from app.api.agents import get_current_agent
|
||||
|
||||
# RBAC 权限装饰器
|
||||
from app.dependencies import require_permission
|
||||
|
||||
from app.services.wecom_service import WecomService
|
||||
from app.services.ws_manager import manager
|
||||
from app.utils.response import AppException, ERR_CONVERSATION_NOT_FOUND, ERR_CONVERSATION_RESOLVED, success_response
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -48,6 +55,7 @@ RECALLABLE_WINDOW_MINUTES = 2
|
||||
# GET /api/conversations/{id}/messages — 获取会话消息列表
|
||||
# --------------------------------------------------------------------------
|
||||
@router.get("/conversations/{conversation_id}/messages")
|
||||
@require_permission("conversation", "read", "all")
|
||||
async def list_messages(
|
||||
conversation_id: str,
|
||||
limit: int = Query(50, ge=1, le=100, description="每页消息数量"),
|
||||
@@ -129,6 +137,7 @@ async def list_messages(
|
||||
# POST /api/conversations/{id}/messages — 坐席发送消息
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/messages")
|
||||
@require_permission("conversation", "create", "all")
|
||||
async def send_message(
|
||||
conversation_id: str,
|
||||
body: MessageCreate,
|
||||
@@ -184,6 +193,7 @@ async def send_message(
|
||||
status="sending", # 初始状态为发送中
|
||||
recallable_until=recallable_until,
|
||||
is_read=True, # 坐席自己发的消息默认已读
|
||||
server_timestamp=int(time.time() * 1000), # [MSG-P0-03] 服务端时间戳(毫秒)
|
||||
)
|
||||
db.add(message)
|
||||
|
||||
@@ -235,6 +245,7 @@ async def send_message(
|
||||
# GET /api/conversations/{id}/messages/poll — 坐席轮询新消息
|
||||
# --------------------------------------------------------------------------
|
||||
@router.get("/conversations/{conversation_id}/messages/poll")
|
||||
@require_permission("conversation", "read", "all")
|
||||
async def poll_messages(
|
||||
conversation_id: str,
|
||||
after_message_id: Optional[str] = Query(None, description="返回此消息ID之后的新消息"),
|
||||
@@ -297,6 +308,7 @@ async def poll_messages(
|
||||
# POST /api/messages/{id}/recall — 撤回消息(2分钟内)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/messages/{message_id}/recall")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def recall_message(
|
||||
message_id: str,
|
||||
agent: Agent = Depends(get_current_agent),
|
||||
@@ -342,8 +354,31 @@ async def recall_message(
|
||||
# 将消息内容置为空,表示已撤回
|
||||
message.content = "[消息已撤回]"
|
||||
message.status = "recalled"
|
||||
message.is_recalled = True # MSG-P1-01: 标记为已撤回
|
||||
await db.flush()
|
||||
|
||||
# MSG-P1-01: 通过 WebSocket 广播撤回事件给所有参与者
|
||||
conv_stmt = select(Conversation).where(Conversation.id == message.conversation_id)
|
||||
conv_result = await db.execute(conv_stmt)
|
||||
conversation = conv_result.scalars().first()
|
||||
if conversation:
|
||||
participant_ids = []
|
||||
if conversation.assigned_agent_id:
|
||||
participant_ids.append(conversation.assigned_agent_id)
|
||||
if conversation.employee_id:
|
||||
participant_ids.append(conversation.employee_id)
|
||||
# 广播撤回事件
|
||||
await manager.broadcast_message_status(
|
||||
conv_id=message.conversation_id,
|
||||
msg_id=message.id,
|
||||
status="recalled",
|
||||
participant_ids=participant_ids,
|
||||
extra={
|
||||
"recall_by": agent.user_id,
|
||||
"recall_at": datetime.now().isoformat(),
|
||||
},
|
||||
)
|
||||
|
||||
return success_response(message="消息撤回成功")
|
||||
|
||||
|
||||
@@ -351,6 +386,7 @@ async def recall_message(
|
||||
# DELETE /api/messages/{id} — 删除消息
|
||||
# --------------------------------------------------------------------------
|
||||
@router.delete("/messages/{message_id}")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def delete_message(
|
||||
message_id: str,
|
||||
agent: Agent = Depends(get_current_agent),
|
||||
@@ -394,6 +430,7 @@ async def delete_message(
|
||||
# POST /api/conversations/{id}/mark-read — 标记已读
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/mark-read")
|
||||
@require_permission("conversation", "update", "own")
|
||||
async def mark_read(
|
||||
conversation_id: str,
|
||||
agent: Agent = Depends(get_current_agent),
|
||||
@@ -557,4 +594,117 @@ async def upload_message_file(
|
||||
"file_size": file_size,
|
||||
"content_type": file.content_type,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# GET /api/conversations/{id}/messages/search — 搜索消息(MSG-P1-04)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.get("/conversations/{conversation_id}/messages/search")
|
||||
async def search_messages(
|
||||
conversation_id: str,
|
||||
keyword: str = Query(..., description="搜索关键词"),
|
||||
limit: int = Query(20, ge=1, le=100, description="返回结果数量限制"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""搜索会话消息(按关键词)。
|
||||
|
||||
使用 LIKE 查询匹配消息内容,支持模糊搜索。
|
||||
|
||||
Args:
|
||||
conversation_id: 会话ID
|
||||
keyword: 搜索关键词
|
||||
limit: 返回结果数量限制
|
||||
db: 数据库会话
|
||||
|
||||
Returns:
|
||||
Dict: 统一响应格式,包含匹配的消息列表
|
||||
"""
|
||||
# 校验会话存在
|
||||
conv_id_str = str(conversation_id)
|
||||
conv_stmt = select(Conversation).where(Conversation.id == conv_id_str)
|
||||
conv_result = await db.execute(conv_stmt)
|
||||
conversation = conv_result.scalars().first()
|
||||
if not conversation:
|
||||
raise ERR_CONVERSATION_NOT_FOUND
|
||||
|
||||
# 构建搜索查询(使用 LIKE 进行模糊匹配)
|
||||
# 排除已撤回的消息
|
||||
search_pattern = f"%{keyword}%"
|
||||
stmt = (
|
||||
select(Message)
|
||||
.where(Message.conversation_id == conv_id_str)
|
||||
.where(Message.is_recalled == False) # 排除已撤回的消息
|
||||
.where(Message.content.ilike(search_pattern)) # 不区分大小写匹配
|
||||
.order_by(Message.created_at.desc()) # 最新消息在前
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
messages = list(result.scalars().all())
|
||||
|
||||
# 转换为响应格式
|
||||
items = [MessageResponse.model_validate(m).model_dump() for m in messages]
|
||||
|
||||
return success_response(
|
||||
data={
|
||||
"items": items,
|
||||
"total": len(items),
|
||||
"keyword": keyword,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# POST /api/conversations/{id}/typing — 发送 typing 事件(MSG-P1-03)
|
||||
# --------------------------------------------------------------------------
|
||||
@router.post("/conversations/{conversation_id}/typing")
|
||||
@require_permission("conversation", "read", "all")
|
||||
async def send_typing_event(
|
||||
conversation_id: str,
|
||||
agent: Agent = Depends(get_current_agent),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""发送 typing 事件,通知对方正在输入。
|
||||
|
||||
通过 WebSocket 广播 typing 事件给会话参与者。
|
||||
|
||||
Args:
|
||||
conversation_id: 会话ID
|
||||
agent: 当前坐席
|
||||
db: 数据库会话
|
||||
|
||||
Returns:
|
||||
Dict: 统一响应格式
|
||||
"""
|
||||
# 校验会话存在
|
||||
conv_id_str = str(conversation_id)
|
||||
conv_stmt = select(Conversation).where(Conversation.id == conv_id_str)
|
||||
conv_result = await db.execute(conv_stmt)
|
||||
conversation = conv_result.scalars().first()
|
||||
if not conversation:
|
||||
raise ERR_CONVERSATION_NOT_FOUND
|
||||
|
||||
# 构建参与者列表
|
||||
participant_ids = []
|
||||
if conversation.assigned_agent_id:
|
||||
participant_ids.append(conversation.assigned_agent_id)
|
||||
if conversation.employee_id:
|
||||
participant_ids.append(conversation.employee_id)
|
||||
|
||||
# 广播 typing 事件(排除发送者本人)
|
||||
payload = {
|
||||
"type": "typing",
|
||||
"conv_id": conv_id_str,
|
||||
"sender_id": agent.user_id,
|
||||
"sender_name": agent.name or "坐席",
|
||||
}
|
||||
|
||||
for pid in participant_ids:
|
||||
if pid != agent.user_id: # 不发给自己
|
||||
if pid in manager.active_connections:
|
||||
await manager.send_to_agent(pid, payload)
|
||||
elif pid in manager.employee_connections:
|
||||
await manager.send_to_employee(pid, payload)
|
||||
|
||||
return success_response(message="typing 事件已发送")
|
||||
+52
-13
@@ -12,16 +12,22 @@ import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from functools import wraps
|
||||
from typing import List, Optional
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
from fastapi import Depends, HTTPException, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
from app.config import settings
|
||||
from app.models.agent import Agent
|
||||
from app.services.token_service import TokenService
|
||||
from app.utils.response import AppException
|
||||
|
||||
# 延迟导入 get_current_agent 以避免循环依赖
|
||||
def _get_current_agent():
|
||||
from app.api.agents import get_current_agent
|
||||
return get_current_agent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# HTTP Bearer 认证方案
|
||||
@@ -324,30 +330,63 @@ def require_permission(
|
||||
def decorator(func):
|
||||
sig = inspect.signature(func)
|
||||
params = list(sig.parameters.values())
|
||||
params.append(
|
||||
inspect.Parameter(
|
||||
'current_user',
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
annotation=UserInfo,
|
||||
default=Depends(get_current_user),
|
||||
param_names = {p.name for p in params}
|
||||
|
||||
# 智能检测参数名:优先使用函数已定义的参数名
|
||||
# 支持 current_user (通用/管理端) 和 current_agent (坐席端)
|
||||
if 'current_agent' in param_names:
|
||||
param_name = 'current_agent'
|
||||
param_annotation = Agent
|
||||
param_default = Depends(_get_current_agent())
|
||||
else:
|
||||
param_name = 'current_user'
|
||||
param_annotation = UserInfo
|
||||
param_default = Depends(get_current_user)
|
||||
|
||||
# 检查是否需要添加参数
|
||||
needs_param = param_name not in param_names
|
||||
if needs_param:
|
||||
params.append(
|
||||
inspect.Parameter(
|
||||
param_name,
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
annotation=param_annotation,
|
||||
default=param_default,
|
||||
)
|
||||
)
|
||||
)
|
||||
new_sig = sig.replace(parameters=params)
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
current_user = kwargs.pop('current_user')
|
||||
# 提取注入的用户/坐席信息
|
||||
current_user = kwargs.pop(param_name)
|
||||
|
||||
# 拉用户所有角色的 permissions
|
||||
# 注: UserInfo.roles 是角色名列表,permissions 是 {role: [perm]} 字典
|
||||
# 首次实现简化: 角色判断 + admin 通配符
|
||||
# 完整实现需要查 DB 拉 permissions,见 rbac_service.check_permission
|
||||
|
||||
user_roles = set(current_user.roles or [])
|
||||
# 支持两种类型:
|
||||
# 1. UserInfo (H5/管理端): 有 roles 属性 (List[str])
|
||||
# 2. Agent (坐席端): 有 role 属性 (str)
|
||||
if hasattr(current_user, 'roles'):
|
||||
user_roles = set(current_user.roles or [])
|
||||
user_id = current_user.employee_id
|
||||
elif hasattr(current_user, 'role'):
|
||||
# Agent 类型:role 是字符串,直接作为角色
|
||||
user_roles = {current_user.role} if current_user.role else set()
|
||||
user_id = current_user.user_id
|
||||
else:
|
||||
# 兼容:没有 roles 或 role 属性的情况
|
||||
user_roles = set()
|
||||
user_id = getattr(current_user, 'user_id', 'unknown')
|
||||
|
||||
# 保存 user_id 供后续使用
|
||||
current_user._rbac_user_id = user_id
|
||||
|
||||
# 1. admin 角色直通(通配符 *:*:all)
|
||||
if "admin" in user_roles:
|
||||
return await func(*args, current_user=current_user, **kwargs)
|
||||
return await func(*args, **{param_name: current_user}, **kwargs)
|
||||
|
||||
# 2. 其他角色: 走 rbac_service.check_permission
|
||||
# 简化: 这里只看角色名,不查 DB(性能考虑)
|
||||
@@ -372,7 +411,7 @@ def require_permission(
|
||||
|
||||
if not has_perm:
|
||||
logger.warning(
|
||||
f"用户 {current_user.employee_id} 权限不足: "
|
||||
f"用户 {user_id} 权限不足: "
|
||||
f"角色 {list(user_roles)}, 缺 {perm_string}"
|
||||
)
|
||||
raise HTTPException(
|
||||
@@ -380,7 +419,7 @@ def require_permission(
|
||||
detail=f"权限不足: 需要 {perm_string}",
|
||||
)
|
||||
|
||||
return await func(*args, current_user=current_user, **kwargs)
|
||||
return await func(*args, **{param_name: current_user}, **kwargs)
|
||||
|
||||
wrapper.__signature__ = new_sig
|
||||
return wrapper
|
||||
|
||||
@@ -52,6 +52,7 @@ ROLE_PERMISSIONS: Dict[str, Set[Tuple[str, str, str]]] = {
|
||||
("conversation", "read", "own"),
|
||||
("conversation", "read", "all"), # 看所有未分配的会话(坐席工作台需要)
|
||||
("conversation", "update", "own"),
|
||||
("conversation", "update", "all"), # 抢单需要能更新其他坐席的会话
|
||||
("conversation", "create", "all"),
|
||||
},
|
||||
|
||||
|
||||
Reference in New Issue
Block a user