A组认证加固: P0兜底+Token刷新+环境检测+OTP+RBRAC落地+P1日志审计+Token撤销 - 全局一致性审查通过

This commit is contained in:
Simon
2026-07-02 19:06:12 +08:00
parent 78f60c6857
commit fc22de7f4d
12 changed files with 1813 additions and 106 deletions
+68
View File
@@ -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",
})
+89 -13
View File
@@ -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"],
+365 -66
View File
@@ -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 中存储的 nextfallback 到 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 已过期,请重新登录")
+75
View File
@@ -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
View File
@@ -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
View File
@@ -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
+1
View File
@@ -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"),
},