ForcePilot/backend/server/utils/oidc_utils.py
DSYZayn 01705aa66b fix(auth): 修复OIDC第三方登录认证集成的潜在风险
- 实现OIDC用户自动创建和部门管理逻辑,支持并发场景下的用户去重处理
- 重构OIDC回调流程,采用一次性code机制提升安全性,避免敏感信息通过URL传递
- 增加对已注销OIDC用户的恢复支持,兼容历史后缀用户ID格式
- 优化用户软删除逻辑,使用user_id和id组合生成哈希避免重名冲突
2026-04-01 22:10:24 +08:00

278 lines
9.3 KiB
Python

"""OIDC 认证工具类"""
import secrets
import time
import urllib.parse
from typing import Any, Optional
import httpx
from yuxi.utils import logger
from server.utils.oidc_config import oidc_config
class OIDCProviderMetadata:
"""OIDC Provider 元数据"""
def __init__(self):
self.authorization_endpoint: Optional[str] = None
self.token_endpoint: Optional[str] = None
self.userinfo_endpoint: Optional[str] = None
self.end_session_endpoint: Optional[str] = None
self._loaded = False
async def load(self, issuer_url: str) -> bool:
"""从 discovery 端点加载元数据"""
if self._loaded:
return True
try:
# 构建 discovery URL
discovery_url = f"{issuer_url.rstrip('/')}/.well-known/openid-configuration"
async with httpx.AsyncClient() as client:
response = await client.get(discovery_url, timeout=30.0)
response.raise_for_status()
metadata = response.json()
self.authorization_endpoint = metadata.get("authorization_endpoint")
self.token_endpoint = metadata.get("token_endpoint")
self.userinfo_endpoint = metadata.get("userinfo_endpoint")
self.end_session_endpoint = metadata.get("end_session_endpoint")
self._loaded = True
logger.info(f"OIDC discovery loaded from {discovery_url}")
return True
except Exception as e:
logger.error(f"Failed to load OIDC discovery: {e}")
return False
class OIDCUtils:
"""OIDC 工具类"""
_metadata: Optional[OIDCProviderMetadata] = None
_state_store: dict[str, dict[str, Any]] = {}
_login_code_store: dict[str, dict[str, Any]] = {}
_state_ttl_seconds = 300
_login_code_ttl_seconds = 60
@classmethod
def _cleanup_expired_state(cls) -> None:
now = time.time()
expired = [k for k, v in cls._state_store.items() if v["expires_at"] <= now]
for key in expired:
cls._state_store.pop(key, None)
@classmethod
def _cleanup_expired_login_code(cls) -> None:
now = time.time()
expired = [k for k, v in cls._login_code_store.items() if v["expires_at"] <= now]
for key in expired:
cls._login_code_store.pop(key, None)
@classmethod
async def get_metadata(cls) -> Optional[OIDCProviderMetadata]:
"""获取 OIDC Provider 元数据"""
if not oidc_config.enabled or not oidc_config.is_configured():
return None
if cls._metadata is None:
cls._metadata = OIDCProviderMetadata()
# 优先使用配置中的端点
if oidc_config.authorization_endpoint:
cls._metadata.authorization_endpoint = oidc_config.authorization_endpoint
cls._metadata.token_endpoint = oidc_config.token_endpoint
cls._metadata.userinfo_endpoint = oidc_config.userinfo_endpoint
cls._metadata.end_session_endpoint = oidc_config.end_session_endpoint
cls._metadata._loaded = True
else:
# 从 discovery 加载
success = await cls._metadata.load(oidc_config.issuer_url)
if not success:
return None
return cls._metadata
@classmethod
def generate_state(cls, redirect_path: str = "/") -> str:
"""生成 state 参数并存储"""
cls._cleanup_expired_state()
state = secrets.token_urlsafe(32)
cls._state_store[state] = {
"redirect_path": redirect_path,
"expires_at": time.time() + cls._state_ttl_seconds,
}
return state
@classmethod
def verify_state(cls, state: str) -> Optional[dict[str, Any]]:
"""验证 state 参数"""
state_data = cls._state_store.pop(state, None)
if not state_data:
return None
if state_data["expires_at"] <= time.time():
return None
return {"redirect_path": state_data["redirect_path"]}
@classmethod
def generate_login_code(cls, payload: dict[str, Any]) -> str:
"""生成一次性短期登录 code"""
cls._cleanup_expired_login_code()
code = secrets.token_urlsafe(32)
cls._login_code_store[code] = {
"payload": payload,
"expires_at": time.time() + cls._login_code_ttl_seconds,
}
return code
@classmethod
def consume_login_code(cls, code: str) -> Optional[dict[str, Any]]:
"""消费一次性短期登录 code"""
data = cls._login_code_store.pop(code, None)
if not data:
return None
if data["expires_at"] <= time.time():
return None
return data["payload"]
@classmethod
def generate_nonce(cls) -> str:
"""生成 nonce 参数"""
return secrets.token_urlsafe(32)
@classmethod
async def build_authorization_url(cls, redirect_path: str = "/") -> Optional[str]:
"""构建授权 URL"""
metadata = await cls.get_metadata()
if not metadata or not metadata.authorization_endpoint:
return None
state = cls.generate_state(redirect_path)
nonce = cls.generate_nonce()
# 构建 redirect_uri
redirect_uri = oidc_config.redirect_uri
if not redirect_uri:
# 自动构建回调 URL
redirect_uri = "/api/auth/oidc/callback"
params = {
"client_id": oidc_config.client_id,
"response_type": "code",
"scope": oidc_config.scopes,
"redirect_uri": redirect_uri,
"state": state,
"nonce": nonce,
}
query_string = urllib.parse.urlencode(params)
return f"{metadata.authorization_endpoint}?{query_string}"
@classmethod
async def exchange_code_for_token(cls, code: str) -> Optional[dict[str, Any]]:
"""用授权码交换令牌"""
metadata = await cls.get_metadata()
if not metadata or not metadata.token_endpoint:
return None
redirect_uri = oidc_config.redirect_uri or "/api/auth/oidc/callback"
data = {
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"client_id": oidc_config.client_id,
"client_secret": oidc_config.client_secret,
}
try:
async with httpx.AsyncClient() as client:
response = await client.post(
metadata.token_endpoint,
data=data,
headers={"Content-Type": "application/x-www-form-urlencoded"},
timeout=30.0
)
response.raise_for_status()
return response.json()
except Exception as e:
logger.error(f"Failed to exchange code for token: {e}")
return None
@classmethod
async def get_userinfo(cls, access_token: str) -> Optional[dict[str, Any]]:
"""获取用户信息"""
metadata = await cls.get_metadata()
if not metadata or not metadata.userinfo_endpoint:
return None
try:
async with httpx.AsyncClient() as client:
response = await client.get(
metadata.userinfo_endpoint,
headers={"Authorization": f"Bearer {access_token}"},
timeout=30.0
)
response.raise_for_status()
return response.json()
except Exception as e:
logger.error(f"Failed to get userinfo: {e}")
return None
@classmethod
async def build_logout_url(cls, id_token: Optional[str] = None) -> Optional[str]:
"""构建登出 URL"""
metadata = await cls.get_metadata()
if not metadata or not metadata.end_session_endpoint:
return None
params = {"client_id": oidc_config.client_id}
if id_token:
params["id_token_hint"] = id_token
if oidc_config.redirect_uri:
params["post_logout_redirect_uri"] = oidc_config.redirect_uri
query_string = urllib.parse.urlencode(params)
return f"{metadata.end_session_endpoint}?{query_string}"
@classmethod
def extract_user_info(cls, userinfo: dict[str, Any]) -> dict[str, Any]:
"""从 userinfo 中提取用户信息"""
# 获取 sub (subject) - OIDC 用户的唯一标识
sub = userinfo.get("sub", "")
# 获取用户名
username = userinfo.get(oidc_config.username_claim, "")
if not username:
username = userinfo.get("preferred_username", "")
if not username:
username = userinfo.get("email", "").split("@")[0]
if not username:
username = sub[:20] # 使用 sub 的前20位
# 获取邮箱
email = userinfo.get(oidc_config.email_claim, "")
if not email:
email = userinfo.get("email", "")
# 获取显示名称
name = userinfo.get(oidc_config.name_claim, "")
if not name:
name = userinfo.get("name", "")
if not name:
name = username
return {
"sub": sub,
"username": username,
"email": email,
"name": name,
"raw": userinfo,
}